From d2c554b69db656b0156c97d3bc976980f44c68a2 Mon Sep 17 00:00:00 2001
From: Zuellni <123005779+Zuellni@users.noreply.github.com>
Date: Wed, 25 Oct 2023 19:53:44 +0200
Subject: [PATCH 1/7] Add loras, 8bit cache, fix random seed, unloading
---
README.md | 5 +-
exllama.py | 219 +++++++++++++++++++++++++++++++----------------
requirements.txt | 3 +-
3 files changed, 152 insertions(+), 75 deletions(-)
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"
From 0b94a076abcae3c06829ebad1a09ac804cbf028f Mon Sep 17 00:00:00 2001
From: Zuellni <123005779+Zuellni@users.noreply.github.com>
Date: Wed, 25 Oct 2023 19:57:16 +0200
Subject: [PATCH 2/7] Update README.md
---
README.md | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/README.md b/README.md
index 8e9dd49..4a839f4 100644
--- a/README.md
+++ b/README.md
@@ -13,7 +13,7 @@ If you see any ExLlama-related errors while loading, install it manually followi
## Nodes
Name | Description
:--- | :---
-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.
+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.
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.
From 2b3ddde76b8160bd7b7b1b03067adcc7c15c5c18 Mon Sep 17 00:00:00 2001
From: Zuellni <123005779+Zuellni@users.noreply.github.com>
Date: Wed, 25 Oct 2023 19:59:58 +0200
Subject: [PATCH 3/7] Update README.md
---
README.md | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/README.md b/README.md
index 4a839f4..0951199 100644
--- a/README.md
+++ b/README.md
@@ -14,7 +14,7 @@ If you see any ExLlama-related errors while loading, install it manually followi
Name | Description
:--- | :---
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.
-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.
+LoRA | Used to load LoRAs. The directory should contain `adapter_model.bin`/`adapter_model.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.
From dcaade09c243a819c3e11dafcc17734d8851d925 Mon Sep 17 00:00:00 2001
From: Zuellni <123005779+Zuellni@users.noreply.github.com>
Date: Fri, 27 Oct 2023 15:51:15 +0200
Subject: [PATCH 4/7] Remove loras for now, there seems to be a memory leak and
idk how to fix it Remove allowed strings, they don't seem very useful
---
README.md | 4 +-
exllama.py | 147 +++++++++--------------------------------------------
text.py | 50 ------------------
3 files changed, 24 insertions(+), 177 deletions(-)
diff --git a/README.md b/README.md
index 0951199..33d6843 100644
--- a/README.md
+++ b/README.md
@@ -13,10 +13,8 @@ If you see any ExLlama-related errors while loading, install it manually followi
## Nodes
Name | Description
:--- | :---
-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.
-LoRA | Used to load LoRAs. The directory should contain `adapter_model.bin`/`adapter_model.safetensors` and `adapter_config.json`. LoRA parameter count has to match the model.
+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.
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 a76ee6b..c9b7c05 100644
--- a/exllama.py
+++ b/exllama.py
@@ -5,17 +5,11 @@ from time import time
import torch
from comfy.model_management import soft_empty_cache
from comfy.utils import ProgressBar
-from exllamav2 import (
- ExLlamaV2,
- ExLlamaV2Cache_8bit,
- ExLlamaV2Config,
- ExLlamaV2Lora,
- ExLlamaV2Tokenizer,
-)
+from exllamav2 import ExLlamaV2, ExLlamaV2Cache_8bit, ExLlamaV2Config, ExLlamaV2Tokenizer
from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator
-class Model:
+class Loader:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -26,7 +20,7 @@ class Model:
}
CATEGORY = "Zuellni/ExLlama"
- FUNCTION = "prepare"
+ FUNCTION = "process"
RETURN_NAMES = ("MODEL",)
RETURN_TYPES = ("EXL_MODEL",)
@@ -37,9 +31,8 @@ class Model:
self.tokenizer = None
self.generator = None
- def prepare(self, model_dir, max_seq_len):
+ def process(self, model_dir, max_seq_len):
self.unload()
-
self.config = ExLlamaV2Config()
self.config.model_dir = model_dir
self.config.prepare()
@@ -47,63 +40,25 @@ class Model:
if max_seq_len:
self.config.max_seq_len = max_seq_len
+ self.tokenizer = ExLlamaV2Tokenizer(self.config)
self.load()
- return ((self, []),)
+ 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
+ self.generator = ExLlamaV2StreamingGenerator(self.base, self.cache, self.tokenizer)
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, lora_dir):
- model, loras = model
-
- lora = ExLlamaV2Lora.from_directory(model.load(), lora_dir)
- loras = loras.copy()
- loras.append(lora)
-
- return ((model, loras),)
+ gc.collect()
+ soft_empty_cache()
class Generator:
@@ -112,6 +67,8 @@ class Generator:
return {
"required": {
"model": ("EXL_MODEL",),
+ "unload": ("BOOLEAN", {"default": False}),
+ "stop_on_newline": ("BOOLEAN", {"default": False}),
"max_new_tokens": ("INT", {"default": 128, "max": 8192}),
"temperature": ("FLOAT", {"default": 0.7, "max": 2, "step": 0.01}),
"top_k": ("INT", {"default": 20, "max": 200}),
@@ -119,10 +76,6 @@ 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}),
- "allow_strings": ("BOOLEAN", {"default": False}),
- "strings": ("STRING", {"default": ""}),
"text": ("STRING", {"multiline": True}),
},
"hidden": {
@@ -136,37 +89,11 @@ 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,
+ unload,
+ stop_on_newline,
max_new_tokens,
temperature,
top_k,
@@ -174,10 +101,6 @@ class Generator:
typical_p,
penalty,
seed,
- unload,
- stop_on_newline,
- allow_strings,
- strings,
text,
info=None,
id=None,
@@ -185,18 +108,19 @@ class Generator:
if not text:
return ("",)
- model, loras = model
-
model.load()
- text = model.tokenizer.encode(text)
+ input = model.tokenizer.encode(text)
stop_conditions = [model.tokenizer.eos_token_id]
- if not max_new_tokens:
- max_new_tokens = model.config.max_seq_len - text.shape[-1]
-
if stop_on_newline:
stop_conditions.append(model.tokenizer.newline_token_id)
+ if not max_new_tokens:
+ max_new_tokens = model.config.max_seq_len - input.shape[-1]
+
+ model.generator.set_stop_conditions(stop_conditions)
+ random.seed(seed)
+
settings = ExLlamaV2Sampler.Settings()
settings.temperature = temperature
settings.top_k = top_k
@@ -204,23 +128,7 @@ class Generator:
settings.typical = typical_p
settings.token_repetition_penalty = penalty
- if strings:
- strings = self.format(strings)
- tokens = model.tokenizer.encode(strings)
- vocab_size = model.config.vocab_size
- padding = vocab_size + (-vocab_size % 32)
-
- 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")
-
- random.seed(seed)
- model.generator.set_stop_conditions(stop_conditions)
- model.generator.begin_stream(text, settings, loras=loras)
+ model.generator.begin_stream(input, settings, token_healing=True)
progress = ProgressBar(max_new_tokens)
start = time()
eos = False
@@ -229,13 +137,6 @@ class Generator:
while not eos and tokens < max_new_tokens:
chunk, eos, _ = model.generator.stream()
-
- if strings and allow_strings:
- c = (output + chunk).strip()
-
- if not any(c in s for s in strings):
- break
-
progress.update(1)
output += chunk
tokens += 1
@@ -259,13 +160,11 @@ class Generator:
NODE_CLASS_MAPPINGS = {
- "ZuellniExLlamaModel": Model,
- "ZuellniExLlamaLora": Lora,
+ "ZuellniExLlamaLoader": Loader,
"ZuellniExLlamaGenerator": Generator,
}
NODE_DISPLAY_NAME_MAPPINGS = {
- "ZuellniExLlamaModel": "Model",
- "ZuellniExLlamaLora": "LoRA",
+ "ZuellniExLlamaLoader": "Loader",
"ZuellniExLlamaGenerator": "Generator",
}
diff --git a/text.py b/text.py
index 43a9497..8a26f62 100644
--- a/text.py
+++ b/text.py
@@ -1,51 +1,3 @@
-from comfy.model_management import InterruptProcessingException
-
-
-class Condition:
- @classmethod
- def INPUT_TYPES(cls):
- return {
- "required": {
- "a": ("STRING", {"forceInput": True}),
- "condition": (["==", "!=", ">", ">=", "<", "<=", "in", "sw", "ew"],),
- "b": ("STRING", {"default": ""}),
- },
- "optional": {
- "text": ("STRING", {"forceInput": True, "multiline": True}),
- },
- }
-
- CATEGORY = "Zuellni/Text"
- FUNCTION = "condition"
- OUTPUT_Node = True
- RETURN_NAMES = ("TEXT",)
- RETURN_TYPES = ("STRING",)
-
- def condition(self, a, condition, b, text=None):
- try:
- a = float(a)
- b = float(b)
- except:
- pass
-
- conditions = {
- "==": lambda: a == b,
- "!=": lambda: a != b,
- ">": lambda: a > b,
- ">=": lambda: a >= b,
- "<": lambda: a < b,
- "<=": lambda: a <= b,
- "in": lambda: str(a) in str(b),
- "sw": lambda: str(a).startswith(str(b)),
- "ew": lambda: str(a).endswith(str(b)),
- }
-
- if not conditions[condition]():
- raise InterruptProcessingException()
-
- return (text,)
-
-
class Format:
@classmethod
def INPUT_TYPES(cls):
@@ -89,13 +41,11 @@ class Preview:
NODE_CLASS_MAPPINGS = {
- "ZuellniTextCondition": Condition,
"ZuellniTextFormat": Format,
"ZuellniTextPreview": Preview,
}
NODE_DISPLAY_NAME_MAPPINGS = {
- "ZuellniTextCondition": "Condition",
"ZuellniTextFormat": "Format",
"ZuellniTextPreview": "Preview",
}
From 55871388b69a9ad007fba76163bcb8bce046f5ce Mon Sep 17 00:00:00 2001
From: Zuellni <123005779+Zuellni@users.noreply.github.com>
Date: Fri, 27 Oct 2023 15:54:24 +0200
Subject: [PATCH 5/7] Move for clarity
---
exllama.py | 6 +++---
1 file changed, 3 insertions(+), 3 deletions(-)
diff --git a/exllama.py b/exllama.py
index c9b7c05..4e2f96e 100644
--- a/exllama.py
+++ b/exllama.py
@@ -112,12 +112,12 @@ class Generator:
input = model.tokenizer.encode(text)
stop_conditions = [model.tokenizer.eos_token_id]
- if stop_on_newline:
- stop_conditions.append(model.tokenizer.newline_token_id)
-
if not max_new_tokens:
max_new_tokens = model.config.max_seq_len - input.shape[-1]
+ if stop_on_newline:
+ stop_conditions.append(model.tokenizer.newline_token_id)
+
model.generator.set_stop_conditions(stop_conditions)
random.seed(seed)
From f48ebd1c68a0ad3f2ce66e42a0d4395ed491f1ca Mon Sep 17 00:00:00 2001
From: Zuellni <123005779+Zuellni@users.noreply.github.com>
Date: Tue, 7 Nov 2023 12:46:50 +0100
Subject: [PATCH 6/7] Some minor changes
---
README.md | 2 +-
exllama.py | 17 ++++++++++++-----
text.js | 4 ++--
text.py | 44 ++++++++++++++++++++++++--------------------
4 files changed, 39 insertions(+), 28 deletions(-)
diff --git a/README.md b/README.md
index 33d6843..c45238c 100644
--- a/README.md
+++ b/README.md
@@ -15,7 +15,7 @@ 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.
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).
-Format | Replaces variables enclosed in brackets, such as `[a]`, with their values.
+Replace | Replaces variables enclosed in brackets, such as `[a]`, with their values.
Preview | Displays generated outputs in the UI.
## Workflow
diff --git a/exllama.py b/exllama.py
index 4e2f96e..33ba998 100644
--- a/exllama.py
+++ b/exllama.py
@@ -3,11 +3,13 @@ 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_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
+
class Loader:
@classmethod
@@ -15,7 +17,7 @@ class Loader:
return {
"required": {
"model_dir": ("STRING", {"default": ""}),
- "max_seq_len": ("INT", {"default": 2048, "max": 8192}),
+ "max_seq_len": ("INT", {"default": 1024, "max": MAX}),
},
}
@@ -50,7 +52,12 @@ class Loader:
self.base = ExLlamaV2(self.config)
self.base.load()
self.cache = ExLlamaV2Cache_8bit(self.base)
- self.generator = ExLlamaV2StreamingGenerator(self.base, self.cache, self.tokenizer)
+
+ self.generator = ExLlamaV2StreamingGenerator(
+ self.base,
+ self.cache,
+ self.tokenizer
+ )
def unload(self):
self.base = None
@@ -69,7 +76,7 @@ class Generator:
"model": ("EXL_MODEL",),
"unload": ("BOOLEAN", {"default": False}),
"stop_on_newline": ("BOOLEAN", {"default": False}),
- "max_new_tokens": ("INT", {"default": 128, "max": 8192}),
+ "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}),
diff --git a/text.js b/text.js
index 828c395..3894b88 100644
--- a/text.js
+++ b/text.js
@@ -14,9 +14,9 @@ app.registerExtension({
const position = this.widgets.findIndex((w) => w.name === "text");
if (position !== -1) {
- for (let i = position; i < this.widgets.length; i++) {
+ for (let i = position; i < this.widgets.length; i++)
this.widgets[i].onRemove?.();
- }
+
this.widgets.length = position;
}
diff --git a/text.py b/text.py
index 8a26f62..6dfc23b 100644
--- a/text.py
+++ b/text.py
@@ -1,4 +1,22 @@
-class Format:
+class Preview:
+ @classmethod
+ def INPUT_TYPES(cls):
+ return {
+ "required": {
+ "text": ("STRING", {"forceInput": True, "multiline": True}),
+ }
+ }
+
+ CATEGORY = "Zuellni/Text"
+ FUNCTION = "preview"
+ OUTPUT_NODE = True
+ RETURN_TYPES = ()
+
+ def preview(self, text):
+ return {"ui": {"text": [text]}}
+
+
+class Replace:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -14,40 +32,26 @@ class Format:
}
CATEGORY = "Zuellni/Text"
- FUNCTION = "format"
+ FUNCTION = "replace"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
- def format(self, text, **vars):
- for key, value in vars.items():
+ def replace(self, text, **inputs):
+ for key, value in inputs.items():
if value:
text = text.replace(f"[{key}]", value)
return (text,)
-class Preview:
- @classmethod
- def INPUT_TYPES(cls):
- return {"required": {"text": ("STRING", {"forceInput": True})}}
-
- CATEGORY = "Zuellni/Text"
- FUNCTION = "preview"
- OUTPUT_NODE = True
- RETURN_TYPES = ()
-
- def preview(self, text):
- return {"ui": {"text": [text]}}
-
-
NODE_CLASS_MAPPINGS = {
- "ZuellniTextFormat": Format,
"ZuellniTextPreview": Preview,
+ "ZuellniTextReplace": Replace,
}
NODE_DISPLAY_NAME_MAPPINGS = {
- "ZuellniTextFormat": "Format",
"ZuellniTextPreview": "Preview",
+ "ZuellniTextReplace": "Replace",
}
WEB_DIRECTORY = "."
From 7e97b26af923bf0ce8cdab011fd6b329422aa962 Mon Sep 17 00:00:00 2001
From: Zuellni <123005779+Zuellni@users.noreply.github.com>
Date: Tue, 7 Nov 2023 12:55:39 +0100
Subject: [PATCH 7/7] Fix preview size
---
text.py | 2 +-
1 file changed, 1 insertion(+), 1 deletion(-)
diff --git a/text.py b/text.py
index 6dfc23b..36439d7 100644
--- a/text.py
+++ b/text.py
@@ -3,7 +3,7 @@ class Preview:
def INPUT_TYPES(cls):
return {
"required": {
- "text": ("STRING", {"forceInput": True, "multiline": True}),
+ "text": ("STRING", {"forceInput": True}),
}
}