Fix typo in js script, add printing to console for preview node, remove deprecated config option

This commit is contained in:
Zuellni
2024-12-06 15:22:11 +01:00
parent 63c7394d2c
commit e98048093c
4 changed files with 11 additions and 23 deletions
-5
View File
@@ -43,11 +43,6 @@ git clone https://huggingface.co/turboderp/Llama-3.1-8B-Instruct-exl2 -b 4.0bpw
<td><i>cache_bits</i></td> <td><i>cache_bits</i></td>
<td>A lower value reduces VRAM usage, but also affects generation speed and quality.</td> <td>A lower value reduces VRAM usage, but also affects generation speed and quality.</td>
</tr> </tr>
<tr>
<td></td>
<td><i>fast_tensors</i></td>
<td>Enabling reduces RAM usage and speeds up model loading.</td>
</tr>
<tr> <tr>
<td></td> <td></td>
<td ><i>flash_attention</i></td> <td ><i>flash_attention</i></td>
+7 -16
View File
@@ -27,12 +27,11 @@ from folder_paths import add_model_folder_path, get_folder_paths, models_dir
_CATEGORY = "zuellni/exllama" _CATEGORY = "zuellni/exllama"
_MAPPING = "ZuellniExLlama" _MAPPING = "ZuellniExLlama"
class Loader:
_input_info = None
class Loader:
@classmethod @classmethod
def INPUT_TYPES(cls): def INPUT_TYPES(cls):
def get_input_info(cls): if not cls._MODELS:
add_model_folder_path("llm", str(Path(models_dir) / "llm")) add_model_folder_path("llm", str(Path(models_dir) / "llm"))
for folder in get_folder_paths("llm"): for folder in get_folder_paths("llm"):
@@ -41,21 +40,14 @@ class Loader:
parent = path.relative_to(folder).parent parent = path.relative_to(folder).parent
cls._MODELS[str(parent / path.name)] = path cls._MODELS[str(parent / path.name)] = path
models = list(cls._MODELS.keys()) models = list(cls._MODELS.keys())
caches = list(cls._CACHES.keys()) caches = list(cls._CACHES.keys())
default = models[0] if models else None default = models[0] if models else None
return models, caches, default
if Loader._input_info is None:
Loader._input_info = get_input_info(cls)
models, caches, default = Loader._input_info
return { return {
"required": { "required": {
"model": (models, {"default": default}), "model": (models, {"default": default}),
"cache_bits": (caches, {"default": 4}), "cache_bits": (caches, {"default": 4}),
"fast_tensors": ("BOOLEAN", {"default": True}),
"flash_attention": ("BOOLEAN", {"default": True}), "flash_attention": ("BOOLEAN", {"default": True}),
"max_seq_len": ( "max_seq_len": (
"INT", "INT",
@@ -76,12 +68,11 @@ class Loader:
RETURN_NAMES = ("MODEL",) RETURN_NAMES = ("MODEL",)
RETURN_TYPES = ("EXL_MODEL",) RETURN_TYPES = ("EXL_MODEL",)
def setup(self, model, cache_bits, fast_tensors, flash_attention, max_seq_len): def setup(self, model, cache_bits, flash_attention, max_seq_len):
self.unload() self.unload()
self.cache_bits = cache_bits self.cache_bits = cache_bits
self.config = ExLlamaV2Config(__class__._MODELS[model]) self.config = ExLlamaV2Config(__class__._MODELS[model])
self.config.fasttensors = fast_tensors
self.config.no_flash_attn = not flash_attention self.config.no_flash_attn = not flash_attention
if max_seq_len: if max_seq_len:
@@ -328,7 +319,7 @@ class Generator:
if not settings: if not settings:
settings = ExLlamaV2Sampler.Settings() settings = ExLlamaV2Sampler.Settings()
settings.greedy() settings = settings.greedy()
job = ExLlamaV2DynamicJob( job = ExLlamaV2DynamicJob(
input_ids=tokens, input_ids=tokens,
+1 -1
View File
@@ -3,7 +3,7 @@ import { app } from "../../../scripts/app.js"
app.registerExtension({ app.registerExtension({
name: "ZuellniText", name: "ZuellniText",
async beforeRegisterNodeDef(nodeType, nodeData, app) { async beforeRegisterNodeDef(nodeType, nodeData, app) {
if (nodeData.category != "Zuellni/Text") if (nodeData.category != "zuellni/text")
return return
const onNodeCreated = nodeType.prototype.onNodeCreated const onNodeCreated = nodeType.prototype.onNodeCreated
+3 -1
View File
@@ -71,6 +71,7 @@ class Preview:
return { return {
"required": { "required": {
"text": ("STRING", {"default": "", "forceInput": True}), "text": ("STRING", {"default": "", "forceInput": True}),
"print_to_console": ("BOOLEAN", {"default": False}),
"output": ("STRING", {"default": "", "multiline": True}), "output": ("STRING", {"default": "", "multiline": True}),
} }
} }
@@ -80,7 +81,8 @@ class Preview:
OUTPUT_NODE = True OUTPUT_NODE = True
RETURN_TYPES = () RETURN_TYPES = ()
def preview(self, text, output): def preview(self, text, print_to_console, output):
print_to_console and print(text)
return {"ui": {"text": [text]}} return {"ui": {"text": [text]}}