From 6c43e3eda3d49f796a779220298e3c039f32d1fe Mon Sep 17 00:00:00 2001 From: morphicschris Date: Thu, 2 Nov 2023 23:00:32 +0000 Subject: [PATCH] Code to extract data from LLM characters --- nodes/__init__.py | 7 ++++++- nodes/checkpoint_names.py | 24 +++++++++++++++++++++ nodes/local_llm.py | 44 ++++++++++++++++++++++++++++++++++++++- 3 files changed, 73 insertions(+), 2 deletions(-) create mode 100644 nodes/checkpoint_names.py diff --git a/nodes/__init__.py b/nodes/__init__.py index 282558b..6a17864 100644 --- a/nodes/__init__.py +++ b/nodes/__init__.py @@ -1,7 +1,8 @@ from .image_glitcher import ImageGlitcher from .color_stylizer import ColorStylizer -from .local_llm import QueryLocalLLM +from .local_llm import QueryLocalLLM, ExtractCharacterInfo from .sdxl_resolution import SdxlResolution, SdxlResolutionToDimensions +from .checkpoint_names import CheckpointNames NODE_CLASS_MAPPINGS = { ImageGlitcher.NAME: ImageGlitcher, @@ -9,6 +10,8 @@ NODE_CLASS_MAPPINGS = { QueryLocalLLM.NAME: QueryLocalLLM, SdxlResolution.NAME: SdxlResolution, SdxlResolutionToDimensions.NAME: SdxlResolutionToDimensions, + CheckpointNames.NAME: CheckpointNames, + ExtractCharacterInfo.NAME: ExtractCharacterInfo, } # A dictionary that contains the friendly/humanly readable titles for the nodes @@ -18,6 +21,8 @@ NODE_DISPLAY_NAME_MAPPINGS = { "QueryLocalLLM": "Query Local LLM", "SdxlResolution": "SDXL Resolution", "SdxlResolutionToDimensions": "SDXL Resolution To Dimensions", + "CheckPointNames": "Checkpoint Names", + "ExtractCharacterInfo": "Extract Character Info" } __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] diff --git a/nodes/checkpoint_names.py b/nodes/checkpoint_names.py new file mode 100644 index 0000000..8650475 --- /dev/null +++ b/nodes/checkpoint_names.py @@ -0,0 +1,24 @@ +import folder_paths + +class CheckpointNames: + NAME = "Checkpoint Names" + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "ckpt_name": ("ckpt_name", folder_paths.get_filename_list("checkpoints"), "CKPT_NAME"), + } + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("name",) + FUNCTION = "getCheckpoints" + OUTPUT_NODE = False + CATEGORY = "CrasH Utils/Loaders" + + def getCheckpoints(self, checkpoint_name): + return checkpoint_name \ No newline at end of file diff --git a/nodes/local_llm.py b/nodes/local_llm.py index 4b5613f..701fb6b 100644 --- a/nodes/local_llm.py +++ b/nodes/local_llm.py @@ -1,5 +1,8 @@ import requests import json +from PIL import Image +import base64 +import re class QueryLocalLLM: NAME = "Query Local LLM" @@ -65,4 +68,43 @@ class QueryLocalLLM: return generatedText else: print(f"Error {response.status_code}: {response.text}") - return None \ No newline at end of file + return None + +class ExtractCharacterInfo: + NAME = "Extract Character Information" + + def __init__(self): + pass + + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "path": ("STRING", { "multiline": False, "default": "" }), + }, + } + + RETURN_TYPES = ("STRING",) + RETURN_NAMES = ("character_data",) + FUNCTION = "parse_png_text" + OUTPUT_NODE = False + CATEGORY = "CrasH Utils/LLM" + + def parse_png_text(self, path): + # Open the image file + with Image.open(path) as img: + # PNG images can have multiple text chunks, get them all + text_chunks = [chunk for chunk in img.text.values()] + + # If there are no text chunks, raise an error + if not text_chunks: + raise ValueError('No text data found in PNG image.') + + # Decode the first text chunk from base64 to utf-8 + try: + # This assumes the text is base64-encoded as in the JS example + decoded_text = base64.b64decode(text_chunks[0]).decode('utf-8') + print(decoded_text) + return decoded_text + except Exception as e: + raise ValueError('Could not decode the text chunk.') from e \ No newline at end of file