Code to extract data from LLM characters

This commit is contained in:
morphicschris
2023-11-02 23:00:32 +00:00
parent d0b5b54f6d
commit 6c43e3eda3
3 changed files with 73 additions and 2 deletions
+6 -1
View File
@@ -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"]
+24
View File
@@ -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
+43 -1
View File
@@ -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
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