Add loras, 8bit cache, fix random seed, unloading
This commit is contained in:
@@ -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
@@ -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
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user