Add gpu split, add 8bit cache toggle,

add min_p, encode specal tokens,
update requirements versions
This commit is contained in:
Zuellni
2023-11-20 21:44:52 +01:00
parent 9d20724d02
commit a33a1ac395
3 changed files with 67 additions and 52 deletions
+56 -37
View File
@@ -3,12 +3,10 @@ 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 nodes import MAX_RESOLUTION as MAX from exllamav2 import ExLlamaV2, ExLlamaV2Cache, ExLlamaV2Cache_8bit, ExLlamaV2Config, ExLlamaV2Tokenizer
from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator
class Loader: class Loader:
@@ -17,7 +15,9 @@ class Loader:
return { return {
"required": { "required": {
"model_dir": ("STRING", {"default": ""}), "model_dir": ("STRING", {"default": ""}),
"max_seq_len": ("INT", {"default": 1024, "max": MAX}), "gpu_split": ("STRING", {"default": ""}),
"cache_8bit": ("BOOLEAN", {"default": False}),
"max_seq_len": ("INT", {"default": 1024, "max": 2**16}),
}, },
} }
@@ -32,36 +32,45 @@ class Loader:
self.cache = None self.cache = None
self.tokenizer = None self.tokenizer = None
self.generator = None self.generator = None
self.gpu_split = None
self.cache_8bit = False
def process(self, model_dir, max_seq_len): def process(self, model_dir, gpu_split, cache_8bit, max_seq_len):
self.unload() self.unload()
self.config = ExLlamaV2Config() self.config = ExLlamaV2Config()
self.config.model_dir = model_dir self.config.model_dir = model_dir
self.config.prepare() self.config.prepare()
if gpu_split:
self.gpu_split = [float(a) for a in gpu_split.split(",")]
if max_seq_len: if max_seq_len:
self.config.max_seq_len = max_seq_len self.config.max_seq_len = max_seq_len
self.tokenizer = ExLlamaV2Tokenizer(self.config) self.cache_8bit = cache_8bit
self.load() self.load()
return (self,) return (self,)
def load(self): def load(self):
if not self.base: if self.base:
self.base = ExLlamaV2(self.config) return
self.base.load()
self.cache = ExLlamaV2Cache_8bit(self.base)
self.generator = ExLlamaV2StreamingGenerator( self.base = ExLlamaV2(self.config)
self.base, self.base.load(gpu_split=self.gpu_split)
self.cache,
self.tokenizer if self.cache_8bit:
) self.cache = ExLlamaV2Cache_8bit(self.base)
else:
self.cache = ExLlamaV2Cache(self.base)
self.tokenizer = ExLlamaV2Tokenizer(self.config)
self.generator = ExLlamaV2StreamingGenerator(self.base, self.cache, self.tokenizer)
def unload(self): def unload(self):
self.base = None self.base = None
self.cache = None self.cache = None
self.tokenizer = None
self.generator = None self.generator = None
gc.collect() gc.collect()
@@ -75,13 +84,14 @@ class Generator:
"required": { "required": {
"model": ("EXL_MODEL",), "model": ("EXL_MODEL",),
"unload": ("BOOLEAN", {"default": False}), "unload": ("BOOLEAN", {"default": False}),
"stop_on_newline": ("BOOLEAN", {"default": False}), "single_line": ("BOOLEAN", {"default": False}),
"max_new_tokens": ("INT", {"default": 128, "max": MAX}), "max_tokens": ("INT", {"default": 128, "max": 2**16}),
"temperature": ("FLOAT", {"default": 0.7, "max": 2, "step": 0.01}), "temperature": ("FLOAT", {"default": 1, "max": 2, "step": 0.01}),
"top_k": ("INT", {"default": 20, "max": 200}), "min_p": ("FLOAT", {"default": 0.1, "max": 1, "step": 0.01}),
"top_p": ("FLOAT", {"default": 0.9, "max": 1, "step": 0.01}), "top_k": ("INT", {"max": 200}),
"typical_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}), "top_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}),
"penalty": ("FLOAT", {"default": 1.15, "min": 1, "max": 2, "step": 0.01}), "typical": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}),
"penalty": ("FLOAT", {"default": 1, "min": 1, "max": 2, "step": 0.01}),
"seed": ("INT", {"max": 2**64 - 1}), "seed": ("INT", {"max": 2**64 - 1}),
"text": ("STRING", {"multiline": True}), "text": ("STRING", {"multiline": True}),
}, },
@@ -100,12 +110,13 @@ class Generator:
self, self,
model, model,
unload, unload,
stop_on_newline, single_line,
max_new_tokens, max_tokens,
temperature, temperature,
min_p,
top_k, top_k,
top_p, top_p,
typical_p, typical,
penalty, penalty,
seed, seed,
text, text,
@@ -116,33 +127,37 @@ class Generator:
return ("",) return ("",)
model.load() model.load()
input = model.tokenizer.encode(text) input = model.tokenizer.encode(text, encode_special_tokens=True)
stop_conditions = [model.tokenizer.eos_token_id] input_len = input.shape[-1]
max_len = model.config.max_seq_len - input_len
stop = [model.tokenizer.eos_token_id]
if not max_new_tokens: if not max_tokens or max_tokens > max_len:
max_new_tokens = model.config.max_seq_len - input.shape[-1] max_tokens = max_len
if stop_on_newline: if single_line:
stop_conditions.append(model.tokenizer.newline_token_id) stop.append(model.tokenizer.newline_token_id)
model.generator.set_stop_conditions(stop_conditions) model.generator.set_stop_conditions(stop)
torch.manual_seed(seed)
random.seed(seed) random.seed(seed)
settings = ExLlamaV2Sampler.Settings() settings = ExLlamaV2Sampler.Settings()
settings.temperature = temperature settings.temperature = temperature
settings.min_p = min_p
settings.top_k = top_k settings.top_k = top_k
settings.top_p = top_p settings.top_p = top_p
settings.typical = typical_p settings.typical = typical
settings.token_repetition_penalty = penalty settings.token_repetition_penalty = penalty
model.generator.begin_stream(input, settings, token_healing=True) model.generator.begin_stream(input, settings, token_healing=True)
progress = ProgressBar(max_new_tokens) progress = ProgressBar(max_tokens)
start = time() start = time()
eos = False eos = False
output = "" output = ""
tokens = 0 tokens = 0
while not eos and tokens < max_new_tokens: while not eos and tokens < max_tokens:
chunk, eos, _ = model.generator.stream() chunk, eos, _ = model.generator.stream()
progress.update(1) progress.update(1)
output += chunk output += chunk
@@ -151,7 +166,11 @@ class Generator:
output = output.strip() output = output.strip()
total = round(time() - start, 2) total = round(time() - start, 2)
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",
f"({input_len} context, {tokens} tokens, {speed}t/s)",
)
if unload: if unload:
model.unload() model.unload()
+2 -2
View File
@@ -1,2 +1,2 @@
exllamav2>=0.0.7; platform_system == "Linux" exllamav2>=0.0.8; 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" https://github.com/turboderp/exllamav2/releases/download/v0.0.8/exllamav2-0.0.8+cu121-cp311-cp311-win_amd64.whl; platform_system == "Windows"
+9 -13
View File
@@ -2,27 +2,23 @@ import { app } from "../../../scripts/app.js";
import { ComfyWidgets } from "../../../scripts/widgets.js"; import { ComfyWidgets } from "../../../scripts/widgets.js";
app.registerExtension({ app.registerExtension({
name: "ZuellniTextPreview", name: "ZuellniText",
async beforeRegisterNodeDef(nodeType, nodeData, app) { async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === "ZuellniTextPreview") { if (nodeData.name === "ZuellniTextPreview") {
const onExecuted = nodeType.prototype.onExecuted; nodeType.prototype.onExecuted = function(message) {
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments);
if (this.widgets) { if (this.widgets) {
const position = this.widgets.findIndex((w) => w.name === "text"); const index = this.widgets.findIndex((w) => w.name === "output");
if (position !== -1) { if (index !== -1) {
for (let i = position; i < this.widgets.length; i++) for (let i = index; i < this.widgets.length; i++)
this.widgets[i].onRemove?.(); this.widgets[i].onRemove?.();
this.widgets.length = position; this.widgets.length = index;
} }
const type = ["STRING", { multiline: true }]; this.widgets.length = 1;
const widget = ComfyWidgets["STRING"](this, "text", type, app).widget; const options = ["STRING", {multiline: true }]
const widget = ComfyWidgets["STRING"](this, "output", options, app).widget;
widget.inputEl.readOnly = true; widget.inputEl.readOnly = true;
widget.inputEl.style.opacity = 0.7; widget.inputEl.style.opacity = 0.7;
widget.value = message.text; widget.value = message.text;