Code to extract data from LLM characters
This commit is contained in:
+6
-1
@@ -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"]
|
||||
|
||||
@@ -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
@@ -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
|
||||
Reference in New Issue
Block a user