@@ -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
@@ -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
@@ -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"
|
||||||
|
|||||||
@@ -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;
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
@@ -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 = "."
|
||||||
|
|||||||
Reference in New Issue
Block a user