diff --git a/README.md b/README.md
index d9b5d31..fb2a462 100644
--- a/README.md
+++ b/README.md
@@ -15,10 +15,11 @@ 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.
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).
ExLlama isn't deterministic, so the outputs may differ even with the same seed.
-Previewer | Displays generated outputs in the UI and appends them to workflow metadata.
-Replacer | Replaces variables enclosed in brackets, such as `[a]`, with their values.
+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.
## Workflow
-The image below can be opened in ComfyUI. The [model](https://huggingface.co/turboderp/Mistral-7B-instruct-exl2/tree/2.5bpw) uses around 3-4GB of VRAM depending on sequence length.
+The image below can be opened in ComfyUI.

diff --git a/exllama.py b/exllama.py
index 961407a..93a0e9c 100644
--- a/exllama.py
+++ b/exllama.py
@@ -14,7 +14,7 @@ class Loader:
return {
"required": {
"model_dir": ("STRING", {"default": ""}),
- "max_seq_len": ("INT", {"default": 2048, "min": 1, "max": 8192}),
+ "max_seq_len": ("INT", {"default": 2048, "max": 8192}),
},
}
@@ -23,28 +23,25 @@ class Loader:
RETURN_NAMES = ("MODEL",)
RETURN_TYPES = ("EXL_MODEL",)
- def __init__(self):
- self.model = None
-
def load(self, model_dir, max_seq_len):
- del self.model
collect()
soft_empty_cache()
config = ExLlamaV2Config()
config.model_dir = model_dir
config.prepare()
- config.max_seq_len = max_seq_len
- self.model = ExLlamaV2(config)
- self.model.load()
+ if max_seq_len:
+ config.max_seq_len = max_seq_len
- cache = ExLlamaV2Cache(self.model)
+ model = ExLlamaV2(config)
+ model.load()
+
+ cache = ExLlamaV2Cache(model)
tokenizer = ExLlamaV2Tokenizer(config)
- generator = ExLlamaV2StreamingGenerator(self.model, cache, tokenizer)
- settings = ExLlamaV2Sampler.Settings()
+ generator = ExLlamaV2StreamingGenerator(model, cache, tokenizer)
- return ((tokenizer, generator, settings),)
+ return ((tokenizer, generator),)
class Generator:
@@ -53,16 +50,21 @@ class Generator:
return {
"required": {
"model": ("EXL_MODEL",),
- "stop_on_newline": ("BOOLEAN", {"default": False}),
- "max_tokens": ("INT", {"default": 128, "min": 1, "max": 8192}),
+ "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}),
"top_p": ("FLOAT", {"default": 0.9, "max": 1, "step": 0.01}),
- "typical": ("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}),
"seed": ("INT", {"max": 2**64 - 1}),
+ "stop_on_newline": ("BOOLEAN", {"default": False}),
+ "allowed_strings": ("STRING", {"default": ""}),
"text": ("STRING", {"multiline": True}),
},
+ "hidden": {
+ "info": "EXTRA_PNGINFO",
+ "id": "UNIQUE_ID",
+ },
}
CATEGORY = "Zuellni/ExLlama"
@@ -73,51 +75,114 @@ class Generator:
def generate(
self,
model,
- stop_on_newline,
- max_tokens,
+ max_new_tokens,
temperature,
top_k,
top_p,
- typical,
+ typical_p,
penalty,
seed,
+ stop_on_newline,
+ allowed_strings,
text,
+ info=None,
+ id=None,
):
+ text = text.strip()
+
if not text:
return ("",)
- tokenizer, generator, settings = model
- progress = ProgressBar(max_tokens)
- prompt = tokenizer.encode(text)
-
+ tokenizer, generator = model
+ text = tokenizer.encode(text)
stop_conditions = [tokenizer.eos_token_id]
- stop_on_newline and stop_conditions.append(tokenizer.newline_token_id)
- generator.set_stop_conditions(stop_conditions)
+ if not max_new_tokens:
+ max_new_tokens = tokenizer.config.max_seq_len - text.shape[-1]
+
+ if stop_on_newline:
+ stop_conditions.append(tokenizer.newline_token_id)
+
+ settings = ExLlamaV2Sampler.Settings()
settings.temperature = temperature
settings.top_k = top_k
settings.top_p = top_p
- settings.typical = typical
+ 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
+ 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.begin_stream(prompt, settings)
+ generator.set_stop_conditions(stop_conditions)
+ generator.begin_stream(text, settings)
+ progress = ProgressBar(max_new_tokens)
start = time()
eos = False
output = ""
tokens = 0
- while not eos and tokens < max_tokens:
+ while not eos and tokens < max_new_tokens:
chunk, eos, _ = generator.stream()
+
+ if allowed_strings:
+ c = (output + chunk).strip()
+
+ if not any(c in s for s in allowed_strings):
+ break
+
progress.update(1)
output += chunk
tokens += 1
+ output = output.strip()
total = round(time() - start, 2)
speed = round(tokens / total, 2)
print(f"Output generated in {total} seconds ({tokens} tokens, {speed} tokens/s)")
- return (output.strip(),)
+ 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)
+
+ if node:
+ node["widgets_values"] = [output]
+
+ return (output,)
NODE_CLASS_MAPPINGS = {
diff --git a/text.js b/text.js
index 79eef99..828c395 100644
--- a/text.js
+++ b/text.js
@@ -2,9 +2,9 @@ import { app } from "../../../scripts/app.js";
import { ComfyWidgets } from "../../../scripts/widgets.js";
app.registerExtension({
- name: "ZuellniTextPreviewer",
+ name: "ZuellniTextPreview",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
- if (nodeData.name === "ZuellniTextPreviewer") {
+ if (nodeData.name === "ZuellniTextPreview") {
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function (message) {
diff --git a/text.py b/text.py
index 6488c82..43a9497 100644
--- a/text.py
+++ b/text.py
@@ -1,33 +1,52 @@
-class Previewer:
+from comfy.model_management import InterruptProcessingException
+
+
+class Condition:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
- "text": ("STRING", {"forceInput": True}),
+ "a": ("STRING", {"forceInput": True}),
+ "condition": (["==", "!=", ">", ">=", "<", "<=", "in", "sw", "ew"],),
+ "b": ("STRING", {"default": ""}),
},
- "hidden": {
- "info": "EXTRA_PNGINFO",
- "id": "UNIQUE_ID",
+ "optional": {
+ "text": ("STRING", {"forceInput": True, "multiline": True}),
},
}
CATEGORY = "Zuellni/Text"
- FUNCTION = "preview"
- OUTPUT_NODE = True
- RETURN_TYPES = ()
+ FUNCTION = "condition"
+ OUTPUT_Node = True
+ RETURN_NAMES = ("TEXT",)
+ RETURN_TYPES = ("STRING",)
- def preview(self, text, info=None, id=None):
- 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)
+ def condition(self, a, condition, b, text=None):
+ try:
+ a = float(a)
+ b = float(b)
+ except:
+ pass
- if node:
- node["widgets_values"] = [text]
+ 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)),
+ }
- return {"ui": {"text": [text]}}
+ if not conditions[condition]():
+ raise InterruptProcessingException()
+
+ return (text,)
-class Replacer:
+class Format:
@classmethod
def INPUT_TYPES(cls):
return {
@@ -43,25 +62,42 @@ class Replacer:
}
CATEGORY = "Zuellni/Text"
- FUNCTION = "replace"
+ FUNCTION = "format"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
- def replace(self, text, **vars):
+ def format(self, text, **vars):
for key, value in vars.items():
- text = text.replace(f"[{key}]", value)
+ if value:
+ text = text.replace(f"[{key}]", value)
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 = {
- "ZuellniTextPreviewer": Previewer,
- "ZuellniTextReplacer": Replacer,
+ "ZuellniTextCondition": Condition,
+ "ZuellniTextFormat": Format,
+ "ZuellniTextPreview": Preview,
}
NODE_DISPLAY_NAME_MAPPINGS = {
- "ZuellniTextPreviewer": "Preview Text",
- "ZuellniTextReplacer": "Replace Text",
+ "ZuellniTextCondition": "Condition",
+ "ZuellniTextFormat": "Format",
+ "ZuellniTextPreview": "Preview",
}
WEB_DIRECTORY = "."