Merge pull request #9 from Zuellni/0.0.7

0.0.7
This commit is contained in:
Zuellni
2023-11-07 13:03:56 +01:00
committed by GitHub
5 changed files with 81 additions and 146 deletions
+3 -4
View File
@@ -13,10 +13,9 @@ If you see any ExLlama-related errors while loading, install it manually followi
## Nodes ## Nodes
Name | Description 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. 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).<br><br>ExLlama isn't deterministic, so the outputs may differ even with the same seed. 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. Replace | Replaces variables enclosed in brackets, such as `[a]`, with their values.
Format | Replaces variables enclosed in brackets, such as `[a]`, with their values.
Preview | Displays generated outputs in the UI. Preview | Displays generated outputs in the UI.
## Workflow ## Workflow
+60 -79
View File
@@ -1,11 +1,14 @@
from gc import collect import gc
import random
from time import time from time import time
import torch import torch
from exllamav2 import ExLlamaV2, ExLlamaV2Cache_8bit, ExLlamaV2Config, ExLlamaV2Tokenizer
from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator
from comfy.model_management import soft_empty_cache from comfy.model_management import soft_empty_cache
from comfy.utils import ProgressBar from comfy.utils import ProgressBar
from exllamav2 import ExLlamaV2, ExLlamaV2Cache, ExLlamaV2Config, ExLlamaV2Tokenizer from nodes import MAX_RESOLUTION as MAX
from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator
class Loader: class Loader:
@@ -14,34 +17,55 @@ class Loader:
return { return {
"required": { "required": {
"model_dir": ("STRING", {"default": ""}), "model_dir": ("STRING", {"default": ""}),
"max_seq_len": ("INT", {"default": 2048, "max": 8192}), "max_seq_len": ("INT", {"default": 1024, "max": MAX}),
}, },
} }
CATEGORY = "Zuellni/ExLlama" CATEGORY = "Zuellni/ExLlama"
FUNCTION = "load" FUNCTION = "process"
RETURN_NAMES = ("MODEL",) RETURN_NAMES = ("MODEL",)
RETURN_TYPES = ("EXL_MODEL",) RETURN_TYPES = ("EXL_MODEL",)
def load(self, model_dir, max_seq_len): def __init__(self):
collect() self.config = None
soft_empty_cache() self.base = None
self.cache = None
self.tokenizer = None
self.generator = None
config = ExLlamaV2Config() def process(self, model_dir, max_seq_len):
config.model_dir = model_dir self.unload()
config.prepare() self.config = ExLlamaV2Config()
self.config.model_dir = model_dir
self.config.prepare()
if max_seq_len: if max_seq_len:
config.max_seq_len = max_seq_len self.config.max_seq_len = max_seq_len
model = ExLlamaV2(config) self.tokenizer = ExLlamaV2Tokenizer(self.config)
model.load() self.load()
cache = ExLlamaV2Cache(model) return (self,)
tokenizer = ExLlamaV2Tokenizer(config)
generator = ExLlamaV2StreamingGenerator(model, cache, tokenizer)
return ((tokenizer, generator),) def load(self):
if not self.base:
self.base = ExLlamaV2(self.config)
self.base.load()
self.cache = ExLlamaV2Cache_8bit(self.base)
self.generator = ExLlamaV2StreamingGenerator(
self.base,
self.cache,
self.tokenizer
)
def unload(self):
self.base = None
self.cache = None
self.generator = None
gc.collect()
soft_empty_cache()
class Generator: class Generator:
@@ -50,15 +74,15 @@ class Generator:
return { return {
"required": { "required": {
"model": ("EXL_MODEL",), "model": ("EXL_MODEL",),
"max_new_tokens": ("INT", {"default": 128, "max": 8192}), "unload": ("BOOLEAN", {"default": False}),
"stop_on_newline": ("BOOLEAN", {"default": False}),
"max_new_tokens": ("INT", {"default": 128, "max": MAX}),
"temperature": ("FLOAT", {"default": 0.7, "max": 2, "step": 0.01}), "temperature": ("FLOAT", {"default": 0.7, "max": 2, "step": 0.01}),
"top_k": ("INT", {"default": 20, "max": 200}), "top_k": ("INT", {"default": 20, "max": 200}),
"top_p": ("FLOAT", {"default": 0.9, "max": 1, "step": 0.01}), "top_p": ("FLOAT", {"default": 0.9, "max": 1, "step": 0.01}),
"typical_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), "typical_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}),
"penalty": ("FLOAT", {"default": 1.15, "min": 1, "max": 2, "step": 0.01}), "penalty": ("FLOAT", {"default": 1.15, "min": 1, "max": 2, "step": 0.01}),
"seed": ("INT", {"max": 2**64 - 1}), "seed": ("INT", {"max": 2**64 - 1}),
"stop_on_newline": ("BOOLEAN", {"default": False}),
"allowed_strings": ("STRING", {"default": ""}),
"text": ("STRING", {"multiline": True}), "text": ("STRING", {"multiline": True}),
}, },
"hidden": { "hidden": {
@@ -75,6 +99,8 @@ class Generator:
def generate( def generate(
self, self,
model, model,
unload,
stop_on_newline,
max_new_tokens, max_new_tokens,
temperature, temperature,
top_k, top_k,
@@ -82,26 +108,25 @@ class Generator:
typical_p, typical_p,
penalty, penalty,
seed, seed,
stop_on_newline,
allowed_strings,
text, text,
info=None, info=None,
id=None, id=None,
): ):
text = text.strip()
if not text: if not text:
return ("",) return ("",)
tokenizer, generator = model model.load()
text = tokenizer.encode(text) input = model.tokenizer.encode(text)
stop_conditions = [tokenizer.eos_token_id] stop_conditions = [model.tokenizer.eos_token_id]
if not max_new_tokens: if not max_new_tokens:
max_new_tokens = tokenizer.config.max_seq_len - text.shape[-1] max_new_tokens = model.config.max_seq_len - input.shape[-1]
if stop_on_newline: if stop_on_newline:
stop_conditions.append(tokenizer.newline_token_id) stop_conditions.append(model.tokenizer.newline_token_id)
model.generator.set_stop_conditions(stop_conditions)
random.seed(seed)
settings = ExLlamaV2Sampler.Settings() settings = ExLlamaV2Sampler.Settings()
settings.temperature = temperature settings.temperature = temperature
@@ -110,47 +135,7 @@ class Generator:
settings.typical = typical_p settings.typical = typical_p
settings.token_repetition_penalty = penalty settings.token_repetition_penalty = penalty
if allowed_strings: model.generator.begin_stream(input, settings, token_healing=True)
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
padding = vocab_size + (-vocab_size % 32)
settings.token_bias = torch.full((padding,), float("-inf"))
settings.token_bias[allowed_tokens] = 0
torch.manual_seed(seed)
generator.set_stop_conditions(stop_conditions)
generator.begin_stream(text, settings)
progress = ProgressBar(max_new_tokens) progress = ProgressBar(max_new_tokens)
start = time() start = time()
eos = False eos = False
@@ -158,14 +143,7 @@ class Generator:
tokens = 0 tokens = 0
while not eos and tokens < max_new_tokens: while not eos and tokens < max_new_tokens:
chunk, eos, _ = generator.stream() chunk, eos, _ = model.generator.stream()
if allowed_strings:
c = (output + chunk).strip()
if not any(c in s for s in allowed_strings):
break
progress.update(1) progress.update(1)
output += chunk output += chunk
tokens += 1 tokens += 1
@@ -175,6 +153,9 @@ class Generator:
speed = round(tokens / total, 2) speed = round(tokens / total, 2)
print(f"Output generated in {total} seconds ({tokens} tokens, {speed} tokens/s)") print(f"Output generated in {total} seconds ({tokens} tokens, {speed} tokens/s)")
if unload:
model.unload()
if id and info and "workflow" in info: if id and info and "workflow" in info:
nodes = info["workflow"]["nodes"] nodes = info["workflow"]["nodes"]
node = next((n for n in nodes if str(n["id"]) == id), None) node = next((n for n in nodes if str(n["id"]) == id), None)
+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"
+2 -2
View File
@@ -14,9 +14,9 @@ app.registerExtension({
const position = this.widgets.findIndex((w) => w.name === "text"); const position = this.widgets.findIndex((w) => w.name === "text");
if (position !== -1) { 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[i].onRemove?.();
}
this.widgets.length = position; this.widgets.length = position;
} }
+14 -60
View File
@@ -1,52 +1,22 @@
from comfy.model_management import InterruptProcessingException class Preview:
class Condition:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
"required": { "required": {
"a": ("STRING", {"forceInput": True}), "text": ("STRING", {"forceInput": True}),
"condition": (["==", "!=", ">", ">=", "<", "<=", "in", "sw", "ew"],), }
"b": ("STRING", {"default": ""}),
},
"optional": {
"text": ("STRING", {"forceInput": True, "multiline": True}),
},
} }
CATEGORY = "Zuellni/Text" CATEGORY = "Zuellni/Text"
FUNCTION = "condition" FUNCTION = "preview"
OUTPUT_Node = True OUTPUT_NODE = True
RETURN_NAMES = ("TEXT",) RETURN_TYPES = ()
RETURN_TYPES = ("STRING",)
def condition(self, a, condition, b, text=None): def preview(self, text):
try: return {"ui": {"text": [text]}}
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: class Replace:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
return { return {
@@ -62,42 +32,26 @@ class Format:
} }
CATEGORY = "Zuellni/Text" CATEGORY = "Zuellni/Text"
FUNCTION = "format" FUNCTION = "replace"
RETURN_NAMES = ("TEXT",) RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",) RETURN_TYPES = ("STRING",)
def format(self, text, **vars): def replace(self, text, **inputs):
for key, value in vars.items(): for key, value in inputs.items():
if value: if value:
text = text.replace(f"[{key}]", value) text = text.replace(f"[{key}]", value)
return (text,) 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 = { NODE_CLASS_MAPPINGS = {
"ZuellniTextCondition": Condition,
"ZuellniTextFormat": Format,
"ZuellniTextPreview": Preview, "ZuellniTextPreview": Preview,
"ZuellniTextReplace": Replace,
} }
NODE_DISPLAY_NAME_MAPPINGS = { NODE_DISPLAY_NAME_MAPPINGS = {
"ZuellniTextCondition": "Condition",
"ZuellniTextFormat": "Format",
"ZuellniTextPreview": "Preview", "ZuellniTextPreview": "Preview",
"ZuellniTextReplace": "Replace",
} }
WEB_DIRECTORY = "." WEB_DIRECTORY = "."