Split nodes into files, allow loading each separately without cloning the repo, some other minor stuff, install exllamav2 pip package by default

This commit is contained in:
Zuellni
2023-10-07 10:56:18 +02:00
parent 796a6b8f25
commit 6e78068c1a
7 changed files with 198 additions and 180 deletions
+1 -4
View File
@@ -8,10 +8,7 @@ git clone https://github.com/Zuellni/ComfyUI-ExLlama-Nodes
pip install -r requirements.txt pip install -r requirements.txt
``` ```
If you see any ExLlama-related errors while loading, install it manually from [here](https://github.com/turboderp/exllamav2/releases/latest). For example, on Windows with Python 3.10 and CUDA 11.7: If you see any ExLlama-related errors while loading, install it manually following the instructions from [here](https://github.com/turboderp/exllamav2#installation).
```
pip install https://github.com/turboderp/exllamav2/releases/download/v0.0.4/exllamav2-0.0.4+cu117-cp310-cp310-win_amd64.whl
```
## Nodes ## Nodes
Name | Description Name | Description
+7 -15
View File
@@ -1,17 +1,9 @@
from .nodes import Generator, Loader, Previewer, Replacer from . import exllama, text
NODE_CLASS_MAPPINGS = {
"ZuellniExLlamaLoader": Loader,
"ZuellniExLlamaGenerator": Generator,
"ZuellniTextPreviewer": Previewer,
"ZuellniTextReplacer": Replacer,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZuellniExLlamaLoader": "ExLlama Loader",
"ZuellniExLlamaGenerator": "ExLlama Generator",
"ZuellniTextPreviewer": "Preview Text",
"ZuellniTextReplacer": "Replace Text",
}
NODE_CLASS_MAPPINGS = {}
NODE_DISPLAY_NAME_MAPPINGS = {}
WEB_DIRECTORY = "." WEB_DIRECTORY = "."
for module in (exllama, text):
NODE_CLASS_MAPPINGS.update(module.NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(module.NODE_DISPLAY_NAME_MAPPINGS)
+122
View File
@@ -0,0 +1,122 @@
from time import time
import torch
from comfy.utils import ProgressBar
from exllamav2 import ExLlamaV2, ExLlamaV2Cache, ExLlamaV2Config, ExLlamaV2Tokenizer
from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator
class Loader:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_dir": ("STRING", {"default": ""}),
"max_seq_len": ("INT", {"default": 2048, "min": 1, "max": 8192}),
},
}
CATEGORY = "Zuellni/ExLlama"
FUNCTION = "load"
RETURN_NAMES = ("MODEL",)
RETURN_TYPES = ("EL_MODEL",)
def load(self, model_dir, max_seq_len):
config = ExLlamaV2Config()
config.model_dir = model_dir
config.prepare()
config.max_seq_len = max_seq_len
model = ExLlamaV2(config)
model.load()
cache = ExLlamaV2Cache(model)
tokenizer = ExLlamaV2Tokenizer(config)
generator = ExLlamaV2StreamingGenerator(model, cache, tokenizer)
settings = ExLlamaV2Sampler.Settings()
return ((tokenizer, generator, settings),)
class Generator:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("EL_MODEL",),
"stop_on_newline": ((False, True),),
"max_tokens": ("INT", {"default": 128, "min": 1, "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}),
"penalty": ("FLOAT", {"default": 1.15, "min": 1, "max": 2, "step": 0.01}),
"seed": ("INT", {"max": 2**64 - 1}),
"text": ("STRING", {"multiline": True}),
},
}
CATEGORY = "Zuellni/ExLlama"
FUNCTION = "generate"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
def generate(
self,
model,
stop_on_newline,
max_tokens,
temperature,
top_k,
top_p,
typical,
penalty,
seed,
text,
):
tokenizer, generator, settings = model
progress = ProgressBar(max_tokens)
if not text:
return ("",)
prompt = 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)
settings.temperature = temperature
settings.top_k = top_k
settings.top_p = top_p
settings.typical = typical
settings.token_repetition_penalty = penalty
torch.manual_seed(seed)
generator.begin_stream(prompt, settings)
start = time()
eos = False
output = ""
tokens = 0
while not eos and tokens < max_tokens:
chunk, eos, _ = generator.stream()
progress.update(1)
output += chunk
tokens += 1
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(),)
NODE_CLASS_MAPPINGS = {
"ZuellniExLlamaLoader": Loader,
"ZuellniExLlamaGenerator": Generator,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZuellniExLlamaLoader": "Loader",
"ZuellniExLlamaGenerator": "Generator",
}
-157
View File
@@ -1,157 +0,0 @@
import torch
from comfy.utils import ProgressBar
from exllamav2 import ExLlamaV2, ExLlamaV2Cache, ExLlamaV2Config, ExLlamaV2Tokenizer
from exllamav2.generator import ExLlamaV2Sampler, ExLlamaV2StreamingGenerator
class Loader:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model_dir": ("STRING", {"default": ""}),
"max_seq_len": ("INT", {"default": 2048, "min": 1, "max": 8192}),
},
}
CATEGORY = "Zuellni/ExLlama"
FUNCTION = "load"
RETURN_NAMES = ("MODEL",)
RETURN_TYPES = ("EXLLAMA_MODEL",)
def load(self, model_dir, max_seq_len):
config = ExLlamaV2Config()
config.model_dir = model_dir
config.prepare()
config.max_seq_len = max_seq_len
model = ExLlamaV2(config)
model.load()
tokenizer = ExLlamaV2Tokenizer(config)
cache = ExLlamaV2Cache(model)
generator = ExLlamaV2StreamingGenerator(model, cache, tokenizer)
return (generator,)
class Generator:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"model": ("EXLLAMA_MODEL",),
"stop_on_newline": ([False, True], {"default": False}),
"max_tokens": ("INT", {"default": 128, "min": 1, "max": 8192}),
"temperature": ("FLOAT", {"default": 0.7, "min": 0, "max": 2, "step": 0.01}),
"top_k": ("INT", {"default": 20, "min": 0, "max": 200}),
"top_p": ("FLOAT", {"default": 0.9, "min": 0, "max": 1, "step": 0.01}),
"typical": ("FLOAT", {"default": 1, "min": 0, "max": 1, "step": 0.01}),
"penalty": ("FLOAT", {"default": 1.15, "min": 1, "max": 2, "step": 0.01}),
"seed": ("INT", {"default": 0, "min": 0, "max": 2**64 - 1}),
"text": ("STRING", {"default": "", "multiline": True}),
},
}
CATEGORY = "Zuellni/ExLlama"
FUNCTION = "generate"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
def generate(
self,
model,
stop_on_newline,
max_tokens,
temperature,
top_k,
top_p,
typical,
penalty,
seed,
text,
):
torch.manual_seed(seed)
progress = ProgressBar(max_tokens)
prompt = model.tokenizer.encode(text)
stop_conditions = [model.tokenizer.eos_token_id]
if stop_on_newline:
stop_conditions += [model.tokenizer.newline_token_id]
settings = ExLlamaV2Sampler.Settings()
settings.temperature = temperature
settings.top_k = top_k
settings.top_p = top_p
settings.typical = typical
settings.token_repetition_penalty = penalty
model.set_stop_conditions(stop_conditions)
model.begin_stream(prompt, settings)
eos = False
tokens = 0
text = ""
while not eos and tokens < max_tokens:
chunk, eos, _ = model.stream()
progress.update(1)
text += chunk
tokens += 1
return (text.strip(),)
class Previewer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"forceInput": True}),
},
"hidden": {
"info": "EXTRA_PNGINFO",
"id": "UNIQUE_ID",
},
}
CATEGORY = "Zuellni/Text"
FUNCTION = "preview"
OUTPUT_NODE = True
RETURN_TYPES = ()
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)
if node:
node["widgets_values"] = [text]
return {"ui": {"text": [text]}}
class Replacer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"default": "", "multiline": True}),
},
"optional": {
"a": ("STRING", {"forceInput": True, "multiline": True}),
"b": ("STRING", {"forceInput": True, "multiline": True}),
"c": ("STRING", {"forceInput": True, "multiline": True}),
"d": ("STRING", {"forceInput": True, "multiline": True}),
}
}
CATEGORY = "Zuellni/Text"
FUNCTION = "replace"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
def replace(self, text, **vars):
for key, value, in vars.items():
text = text.replace(f"[{key}]", value)
return (text,)
+1 -4
View File
@@ -1,4 +1 @@
https://github.com/turboderp/exllamav2/releases/download/v0.0.5/exllamav2-0.0.5+cu121-cp311-cp311-win_amd64.whl; platform_system == "Windows" and python_version == "3.11" exllamav2
https://github.com/turboderp/exllamav2/releases/download/v0.0.5/exllamav2-0.0.5+cu118-cp310-cp310-win_amd64.whl; platform_system == "Windows" and python_version == "3.10"
https://github.com/turboderp/exllamav2/releases/download/v0.0.5/exllamav2-0.0.5+cu121-cp311-cp311-linux_x86_64.whl; platform_system == "Linux" and python_version == "3.11"
https://github.com/turboderp/exllamav2/releases/download/v0.0.5/exllamav2-0.0.5+cu118-cp310-cp310-linux_x86_64.whl; platform_system == "Linux" and python_version == "3.10"
View File
+67
View File
@@ -0,0 +1,67 @@
class Previewer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"forceInput": True}),
},
"hidden": {
"info": "EXTRA_PNGINFO",
"id": "UNIQUE_ID",
},
}
CATEGORY = "Zuellni/Text"
FUNCTION = "preview"
OUTPUT_NODE = True
RETURN_TYPES = ()
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)
if node:
node["widgets_values"] = [text]
return {"ui": {"text": [text]}}
class Replacer:
@classmethod
def INPUT_TYPES(cls):
return {
"required": {
"text": ("STRING", {"multiline": True}),
},
"optional": {
"a": ("STRING", {"forceInput": True, "multiline": True}),
"b": ("STRING", {"forceInput": True, "multiline": True}),
"c": ("STRING", {"forceInput": True, "multiline": True}),
"d": ("STRING", {"forceInput": True, "multiline": True}),
},
}
CATEGORY = "Zuellni/Text"
FUNCTION = "replace"
RETURN_NAMES = ("TEXT",)
RETURN_TYPES = ("STRING",)
def replace(self, text, **vars):
for key, value in vars.items():
text = text.replace(f"[{key}]", value)
return (text,)
NODE_CLASS_MAPPINGS = {
"ZuellniTextPreviewer": Previewer,
"ZuellniTextReplacer": Replacer,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"ZuellniTextPreviewer": "Preview Text",
"ZuellniTextReplacer": "Replace Text",
}
WEB_DIRECTORY = "."