Add gpu split, add 8bit cache toggle,
add min_p, encode specal tokens, update requirements versions
This commit is contained in:
+56
-37
@@ -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
@@ -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"
|
||||||
|
|||||||
@@ -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;
|
||||||
|
|||||||
Reference in New Issue
Block a user