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
This commit is contained in:
Zuellni
2023-10-27 15:51:15 +02:00
parent 2b3ddde76b
commit dcaade09c2
3 changed files with 24 additions and 177 deletions
+23 -124
View File
@@ -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",
}