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
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.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:
@@ -17,7 +15,9 @@ class Loader:
return {
"required": {
"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.tokenizer = 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.config = ExLlamaV2Config()
self.config.model_dir = model_dir
self.config.prepare()
if gpu_split:
self.gpu_split = [float(a) for a in gpu_split.split(",")]
if max_seq_len:
self.config.max_seq_len = max_seq_len
self.tokenizer = ExLlamaV2Tokenizer(self.config)
self.cache_8bit = cache_8bit
self.load()
return (self,)
def load(self):
if not self.base:
self.base = ExLlamaV2(self.config)
self.base.load()
self.cache = ExLlamaV2Cache_8bit(self.base)
if self.base:
return
self.generator = ExLlamaV2StreamingGenerator(
self.base,
self.cache,
self.tokenizer
)
self.base = ExLlamaV2(self.config)
self.base.load(gpu_split=self.gpu_split)
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):
self.base = None
self.cache = None
self.tokenizer = None
self.generator = None
gc.collect()
@@ -75,13 +84,14 @@ class Generator:
"required": {
"model": ("EXL_MODEL",),
"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}),
"top_k": ("INT", {"default": 20, "max": 200}),
"top_p": ("FLOAT", {"default": 0.9, "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}),
"single_line": ("BOOLEAN", {"default": False}),
"max_tokens": ("INT", {"default": 128, "max": 2**16}),
"temperature": ("FLOAT", {"default": 1, "max": 2, "step": 0.01}),
"min_p": ("FLOAT", {"default": 0.1, "max": 1, "step": 0.01}),
"top_k": ("INT", {"max": 200}),
"top_p": ("FLOAT", {"default": 1, "max": 1, "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}),
"text": ("STRING", {"multiline": True}),
},
@@ -100,12 +110,13 @@ class Generator:
self,
model,
unload,
stop_on_newline,
max_new_tokens,
single_line,
max_tokens,
temperature,
min_p,
top_k,
top_p,
typical_p,
typical,
penalty,
seed,
text,
@@ -116,33 +127,37 @@ class Generator:
return ("",)
model.load()
input = model.tokenizer.encode(text)
stop_conditions = [model.tokenizer.eos_token_id]
input = model.tokenizer.encode(text, encode_special_tokens=True)
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:
max_new_tokens = model.config.max_seq_len - input.shape[-1]
if not max_tokens or max_tokens > max_len:
max_tokens = max_len
if stop_on_newline:
stop_conditions.append(model.tokenizer.newline_token_id)
if single_line:
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)
settings = ExLlamaV2Sampler.Settings()
settings.temperature = temperature
settings.min_p = min_p
settings.top_k = top_k
settings.top_p = top_p
settings.typical = typical_p
settings.typical = typical
settings.token_repetition_penalty = penalty
model.generator.begin_stream(input, settings, token_healing=True)
progress = ProgressBar(max_new_tokens)
progress = ProgressBar(max_tokens)
start = time()
eos = False
output = ""
tokens = 0
while not eos and tokens < max_new_tokens:
while not eos and tokens < max_tokens:
chunk, eos, _ = model.generator.stream()
progress.update(1)
output += chunk
@@ -151,7 +166,11 @@ class Generator:
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)")
print(
f"Output generated in {total} seconds",
f"({input_len} context, {tokens} tokens, {speed}t/s)",
)
if unload:
model.unload()
+2 -2
View File
@@ -1,2 +1,2 @@
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"
exllamav2>=0.0.8; platform_system == "Linux"
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";
app.registerExtension({
name: "ZuellniTextPreview",
name: "ZuellniText",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === "ZuellniTextPreview") {
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function (message) {
onExecuted?.apply(this, arguments);
nodeType.prototype.onExecuted = function(message) {
if (this.widgets) {
const position = this.widgets.findIndex((w) => w.name === "text");
const index = this.widgets.findIndex((w) => w.name === "output");
if (position !== -1) {
for (let i = position; i < this.widgets.length; i++)
if (index !== -1) {
for (let i = index; i < this.widgets.length; i++)
this.widgets[i].onRemove?.();
this.widgets.length = position;
this.widgets.length = index;
}
const type = ["STRING", { multiline: true }];
const widget = ComfyWidgets["STRING"](this, "text", type, app).widget;
this.widgets.length = 1;
const options = ["STRING", {multiline: true }]
const widget = ComfyWidgets["STRING"](this, "output", options, app).widget;
widget.inputEl.readOnly = true;
widget.inputEl.style.opacity = 0.7;
widget.value = message.text;