Add loras, 8bit cache, fix random seed, unloading

This commit is contained in:
Zuellni
2023-10-25 19:53:44 +02:00
parent 420b1fbd2d
commit d2c554b69d
3 changed files with 152 additions and 75 deletions
+3 -2
View File
@@ -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.<br><br>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).<br><br>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.<br><br>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.
+147 -72
View File
@@ -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",
}
+2 -1
View File
@@ -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"