Bump minimum exllamav2 version, add autosplit and q4 cache, some reformatting

This commit is contained in:
Zuellni
2024-04-07 19:35:32 +02:00
parent d2b4d7a796
commit baed69aba5
7 changed files with 57 additions and 61 deletions
+40 -37
View File
@@ -10,12 +10,16 @@ from exllamav2 import (
ExLlamaV2,
ExLlamaV2Cache,
ExLlamaV2Cache_8bit,
ExLlamaV2Cache_Q4,
ExLlamaV2Config,
ExLlamaV2Tokenizer,
)
from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator
from folder_paths import add_model_folder_path, get_folder_paths, models_dir
_CATEGORY = "Zuellni/ExLlama"
_MAPPING = "ZuellniExLlama"
class Loader:
@classmethod
@@ -34,59 +38,57 @@ class Loader:
return {
"required": {
"model": (models, {"default": default}),
"gpu_split": ("STRING", {"default": ""}),
"cache_8bit": ("BOOLEAN", {"default": False}),
"max_seq_len": ("INT", {"default": 1024, "max": 2**16}),
"cache_bits": ((4, 8, 16), {"default": 4}),
"max_seq_len": ("INT", {"default": 2048, "max": 2**20}),
},
}
_MODELS = {}
CATEGORY = "Zuellni/ExLlama"
CATEGORY = _CATEGORY
FUNCTION = "setup"
RETURN_NAMES = ("MODEL",)
RETURN_TYPES = ("EXL_MODEL",)
def setup(self, model, gpu_split, cache_8bit, max_seq_len):
def setup(self, model, cache_bits, max_seq_len):
self.unload()
self.cache_bits = cache_bits
self.config = ExLlamaV2Config()
self.config.model_dir = __class__._MODELS[model]
self.config.prepare()
if max_seq_len:
self.config.max_seq_len = max_seq_len
self.gpu_split = [float(a) for a in gpu_split.split(",") if gpu_split]
self.cache_8bit = cache_8bit
self.config.max_input_len = max_seq_len
self.config.max_attention_len = max_seq_len**2
return (self,)
def load(self):
if (
hasattr(self, "model") and
hasattr(self, "cache") and
hasattr(self, "tokenizer") and
hasattr(self, "generator") and
self.model and
self.cache and
self.tokenizer and
self.generator
hasattr(self, "model")
and hasattr(self, "cache")
and hasattr(self, "tokenizer")
and hasattr(self, "generator")
and self.model
and self.cache
and self.tokenizer
and self.generator
):
return
self.model = ExLlamaV2(self.config)
progress = ProgressBar(len(self.model.modules))
self.model.load(
gpu_split=self.gpu_split,
callback=lambda s, _: progress.update_absolute(s),
)
progress = ProgressBar(len(self.model.modules) + 1)
self.cache = (
ExLlamaV2Cache_8bit(self.model)
if self.cache_8bit
else ExLlamaV2Cache(self.model)
ExLlamaV2Cache_Q4(self.model, lazy=True)
if self.cache_bits == 4
else ExLlamaV2Cache_8bit(self.model, lazy=True)
if self.cache_bits == 8
else ExLlamaV2Cache(self.model, lazy=True)
)
self.model.load_autosplit(self.cache, callback=lambda _, __: progress.update(1))
self.tokenizer = ExLlamaV2Tokenizer(self.config)
self.generator = ExLlamaV2StreamingGenerator(
@@ -116,14 +118,14 @@ class Generator:
"model": ("EXL_MODEL",),
"unload": ("BOOLEAN", {"default": False}),
"single_line": ("BOOLEAN", {"default": False}),
"max_tokens": ("INT", {"default": 128, "max": 2**16}),
"max_tokens": ("INT", {"default": 128, "max": 2**20}),
"temperature": ("FLOAT", {"default": 1, "max": 5, "step": 0.01}),
"top_k": ("INT", {"max": 200}),
"top_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}),
"typical_p": ("FLOAT", {"default": 1, "max": 1, "step": 0.01}),
"min_p": ("FLOAT", {"max": 1, "step": 0.01}),
"top_a": ("FLOAT", {"max": 1, "step": 0.01}),
"penalty": ("FLOAT", {"default": 1, "min": 1, "max": 3, "step": 0.01}),
"repetition_penalty": ("FLOAT", {"default": 1, "min": 1, "max": 3, "step": 0.01}),
"temperature_last": ("BOOLEAN", {"default": True}),
"seed": ("INT", {"max": 2**64 - 1}),
"text": ("STRING", {"multiline": True}),
@@ -134,7 +136,7 @@ class Generator:
},
}
CATEGORY = "Zuellni/ExLlama"
CATEGORY = _CATEGORY
FUNCTION = "generate"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
@@ -151,7 +153,7 @@ class Generator:
typical_p,
min_p,
top_a,
penalty,
repetition_penalty,
temperature_last,
seed,
text,
@@ -187,20 +189,21 @@ class Generator:
settings.typical = typical_p
settings.min_p = min_p
settings.top_a = top_a
settings.token_repetition_penalty = penalty
settings.token_repetition_penalty = repetition_penalty
settings.temperature_last = temperature_last
start = time()
model.generator.begin_stream(input, settings)
model.generator.begin_stream_ex(input, settings)
progress = ProgressBar(max_tokens)
eos = False
output = ""
tokens = 0
while not eos and tokens < max_tokens:
chunk, eos, _ = model.generator.stream()
response = model.generator.stream_ex()
output += response["chunk"]
eos = response["eos"]
progress.update(1)
output += chunk
tokens += 1
output = output.strip()
@@ -226,11 +229,11 @@ class Generator:
NODE_CLASS_MAPPINGS = {
"ZuellniExLlamaLoader": Loader,
"ZuellniExLlamaGenerator": Generator,
f"{_MAPPING}Loader": Loader,
f"{_MAPPING}Generator": Generator,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZuellniExLlamaLoader": "Loader",
"ZuellniExLlamaGenerator": "Generator",
f"{_MAPPING}Loader": "Loader",
f"{_MAPPING}Generator": "Generator",
}
-2
View File
@@ -1,2 +0,0 @@
exllamav2
flash-attn
-2
View File
@@ -1,2 +0,0 @@
https://github.com/turboderp/exllamav2/releases/download/v0.0.12/exllamav2-0.0.12+cu121-cp311-cp311-win_amd64.whl
https://github.com/bdashore3/flash-attention/releases/download/v2.4.2/flash_attn-2.4.2+cu122torch2.1.2cxx11abiFALSE-cp311-cp311-win_amd64.whl
-2
View File
@@ -1,2 +0,0 @@
https://github.com/turboderp/exllamav2/releases/download/v0.0.14/exllamav2-0.0.14+cu121-cp311-cp311-win_amd64.whl
https://github.com/bdashore3/flash-attention/releases/download/v2.5.2/flash_attn-2.5.2+cu122torch2.2.0cxx11abiFALSE-cp311-cp311-win_amd64.whl
+1
View File
@@ -0,0 +1 @@
exllamav2>=0.0.16
+2 -2
View File
@@ -4,7 +4,7 @@ import { ComfyWidgets } from "../../../scripts/widgets.js";
app.registerExtension({
name: "ZuellniText",
async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.name === "ZuellniTextPreview") {
if (nodeData.name === "ZuellniTextPreviewer") {
const onExecuted = nodeType.prototype.onExecuted;
nodeType.prototype.onExecuted = function(message) {
@@ -20,7 +20,7 @@ app.registerExtension({
this.widgets.length = index;
}
const options = ["STRING", {multiline: true }]
const options = ["STRING", {multiline: true }];
const widget = ComfyWidgets["STRING"](this, "output", options, app).widget;
widget.inputEl.readOnly = true;
+14 -16
View File
@@ -1,13 +1,13 @@
class Preview:
_CATEGORY = "Zuellni/Text"
_MAPPING = "ZuellniText"
class Previewer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"forceInput": True}),
}
}
return {"required": {"text": ("STRING", {"forceInput": True})}}
CATEGORY = "Zuellni/Text"
CATEGORY = _CATEGORY
FUNCTION = "preview"
OUTPUT_NODE = True
RETURN_TYPES = ()
@@ -16,13 +16,11 @@ class Preview:
return {"ui": {"text": [text]}}
class Replace:
class Replacer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
},
"required": {"text": ("STRING", {"multiline": True})},
"optional": {
"a": ("STRING", {"forceInput": True, "multiline": True}),
"b": ("STRING", {"forceInput": True, "multiline": True}),
@@ -31,7 +29,7 @@ class Replace:
},
}
CATEGORY = "Zuellni/Text"
CATEGORY = _CATEGORY
FUNCTION = "replace"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
@@ -45,13 +43,13 @@ class Replace:
NODE_CLASS_MAPPINGS = {
"ZuellniTextPreview": Preview,
"ZuellniTextReplace": Replace,
f"{_MAPPING}Previewer": Previewer,
f"{_MAPPING}Replacer": Replacer,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZuellniTextPreview": "Preview",
"ZuellniTextReplace": "Replace",
f"{_MAPPING}Previewer": "Preview",
f"{_MAPPING}Replacer": "Replace",
}
WEB_DIRECTORY = "."