609 lines
24 KiB
Python
609 lines
24 KiB
Python
"""
|
|
DJZ Prompt Extractor - ComfyUI Custom Node
|
|
Extracts text prompts from PNG image metadata generated by ComfyUI.
|
|
|
|
This node mirrors the logic from the ComfyUI-Prompt-Extractor-JSX component,
|
|
allowing images to be passed through the node to extract embedded prompt text.
|
|
|
|
Author: Drift Johnson
|
|
Repository: https://github.com/MushroomFleet/ComfyUI-Prompt-Extractor-JSX
|
|
"""
|
|
|
|
import json
|
|
import os
|
|
from PIL import Image
|
|
import numpy as np
|
|
import torch
|
|
import folder_paths
|
|
|
|
|
|
class DJZ_PromptExtractor:
|
|
"""
|
|
A ComfyUI node that extracts the positive text prompt from PNG metadata.
|
|
|
|
ComfyUI embeds workflow and prompt data in PNG files using tEXt chunks.
|
|
This node parses that metadata, finds all string values, and returns
|
|
the longest one (which is typically the positive prompt).
|
|
|
|
Uses a file selector dropdown like Load Image for easy file selection.
|
|
"""
|
|
|
|
def __init__(self):
|
|
pass
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
input_dir = folder_paths.get_input_directory()
|
|
files = [f for f in os.listdir(input_dir) if os.path.isfile(os.path.join(input_dir, f))]
|
|
# Filter for image files (primarily PNG but allow others)
|
|
image_extensions = ('.png', '.jpg', '.jpeg', '.webp', '.gif', '.bmp')
|
|
files = sorted([f for f in files if f.lower().endswith(image_extensions)])
|
|
|
|
return {
|
|
"required": {
|
|
"image": (files, {"image_upload": True}),
|
|
},
|
|
"optional": {
|
|
"extraction_mode": (["longest_string", "positive_only", "negative_only", "all_prompts"],),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("STRING", "STRING", "STRING", "IMAGE",)
|
|
RETURN_NAMES = ("prompt", "negative_prompt", "metadata_json", "image",)
|
|
FUNCTION = "extract_prompt"
|
|
CATEGORY = "DJZ-Nodes"
|
|
OUTPUT_NODE = False
|
|
|
|
@classmethod
|
|
def IS_CHANGED(cls, image, extraction_mode="longest_string"):
|
|
image_path = folder_paths.get_annotated_filepath(image)
|
|
return os.path.getmtime(image_path)
|
|
|
|
@classmethod
|
|
def VALIDATE_INPUTS(cls, image, extraction_mode="longest_string"):
|
|
if not folder_paths.exists_annotated_filepath(image):
|
|
return f"Invalid image file: {image}"
|
|
return True
|
|
|
|
@staticmethod
|
|
def find_all_strings(obj, strings=None):
|
|
"""Recursively find all strings in a JSON object/array."""
|
|
if strings is None:
|
|
strings = []
|
|
|
|
if isinstance(obj, str):
|
|
strings.append(obj)
|
|
elif isinstance(obj, list):
|
|
for item in obj:
|
|
DJZ_PromptExtractor.find_all_strings(item, strings)
|
|
elif isinstance(obj, dict):
|
|
for value in obj.values():
|
|
DJZ_PromptExtractor.find_all_strings(value, strings)
|
|
|
|
return strings
|
|
|
|
@staticmethod
|
|
def find_strings_by_key(obj, target_keys, strings=None, current_key=""):
|
|
"""Find strings associated with specific keys (e.g., 'text' in CLIPTextEncode)."""
|
|
if strings is None:
|
|
strings = []
|
|
|
|
if isinstance(obj, dict):
|
|
for key, value in obj.items():
|
|
if key in target_keys and isinstance(value, str):
|
|
strings.append(value)
|
|
else:
|
|
DJZ_PromptExtractor.find_strings_by_key(value, target_keys, strings, key)
|
|
elif isinstance(obj, list):
|
|
for item in obj:
|
|
DJZ_PromptExtractor.find_strings_by_key(item, target_keys, strings, current_key)
|
|
|
|
return strings
|
|
|
|
@staticmethod
|
|
def find_prompts_by_node_type(metadata, node_type="CLIPTextEncode"):
|
|
"""Find all text inputs for a specific node type."""
|
|
prompts = []
|
|
|
|
for key in ["prompt", "workflow"]:
|
|
if key not in metadata:
|
|
continue
|
|
|
|
data = metadata[key]
|
|
if isinstance(data, dict):
|
|
for node_id, node_data in data.items():
|
|
if isinstance(node_data, dict):
|
|
class_type = node_data.get("class_type", "")
|
|
if class_type == node_type:
|
|
inputs = node_data.get("inputs", {})
|
|
if isinstance(inputs, dict) and "text" in inputs:
|
|
text = inputs["text"]
|
|
if isinstance(text, str):
|
|
prompts.append(text)
|
|
|
|
# Also check nested 'nodes' array in workflow format
|
|
if "nodes" in node_data and isinstance(node_data["nodes"], list):
|
|
for node in node_data["nodes"]:
|
|
if isinstance(node, dict) and node.get("type") == node_type:
|
|
widgets = node.get("widgets_values", [])
|
|
if widgets and isinstance(widgets, list):
|
|
for val in widgets:
|
|
if isinstance(val, str) and len(val) > 10:
|
|
prompts.append(val)
|
|
|
|
# Also check if workflow has a 'nodes' array at the top level
|
|
if "workflow" in metadata and isinstance(metadata["workflow"], dict):
|
|
nodes = metadata["workflow"].get("nodes", [])
|
|
if isinstance(nodes, list):
|
|
for node in nodes:
|
|
if isinstance(node, dict) and node.get("type") == node_type:
|
|
widgets = node.get("widgets_values", [])
|
|
if widgets and isinstance(widgets, list):
|
|
for val in widgets:
|
|
if isinstance(val, str) and len(val) > 10:
|
|
prompts.append(val)
|
|
|
|
return prompts
|
|
|
|
@staticmethod
|
|
def classify_prompt(text, metadata):
|
|
"""
|
|
Attempt to classify if a prompt is positive or negative based on context.
|
|
Returns 'positive', 'negative', or 'unknown'.
|
|
"""
|
|
text_lower = text.lower()
|
|
|
|
# Common negative prompt indicators
|
|
negative_indicators = [
|
|
"ugly", "blurry", "bad quality", "low quality", "worst quality",
|
|
"deformed", "mutated", "disfigured", "bad anatomy", "watermark",
|
|
"signature", "text", "error", "cropped", "jpeg artifacts",
|
|
"duplicate", "morbid", "mutilated", "poorly drawn"
|
|
]
|
|
|
|
# Count negative indicators
|
|
negative_count = sum(1 for indicator in negative_indicators if indicator in text_lower)
|
|
|
|
# If more than 2 negative indicators, likely a negative prompt
|
|
if negative_count >= 2:
|
|
return "negative"
|
|
|
|
# Check metadata structure for node connections
|
|
for key in ["prompt", "workflow"]:
|
|
if key not in metadata:
|
|
continue
|
|
data = metadata[key]
|
|
if isinstance(data, dict):
|
|
for node_id, node_data in data.items():
|
|
if isinstance(node_data, dict):
|
|
inputs = node_data.get("inputs", {})
|
|
if isinstance(inputs, dict) and inputs.get("text") == text:
|
|
# Check if this node connects to a "negative" input
|
|
for other_id, other_data in data.items():
|
|
if isinstance(other_data, dict):
|
|
other_inputs = other_data.get("inputs", {})
|
|
if isinstance(other_inputs, dict):
|
|
neg = other_inputs.get("negative", [])
|
|
if isinstance(neg, list) and len(neg) >= 1:
|
|
if str(neg[0]) == str(node_id):
|
|
return "negative"
|
|
pos = other_inputs.get("positive", [])
|
|
if isinstance(pos, list) and len(pos) >= 1:
|
|
if str(pos[0]) == str(node_id):
|
|
return "positive"
|
|
|
|
return "unknown"
|
|
|
|
def extract_from_path(self, image_path):
|
|
"""Extract metadata from a PNG file path."""
|
|
if not os.path.exists(image_path):
|
|
return None, f"File not found: {image_path}"
|
|
|
|
try:
|
|
img = Image.open(image_path)
|
|
return img.info, None
|
|
except Exception as e:
|
|
return None, f"Error reading image: {str(e)}"
|
|
|
|
def extract_from_tensor(self, image_tensor):
|
|
"""
|
|
Attempt to extract metadata from image tensor.
|
|
Note: Tensors typically don't preserve PNG metadata, so this may return empty.
|
|
"""
|
|
# ComfyUI image tensors are [B, H, W, C] format
|
|
# PNG metadata is lost when converted to tensor
|
|
# This is a limitation - for full metadata, use image_path
|
|
return {}, "Note: PNG metadata is not preserved in tensor format. Use image_path for full extraction."
|
|
|
|
def extract_prompt(self, image, extraction_mode="longest_string"):
|
|
"""
|
|
Main extraction function.
|
|
|
|
Args:
|
|
image: Filename of the image in the input directory
|
|
extraction_mode: How to extract prompts
|
|
- longest_string: Return the longest string found (default)
|
|
- positive_only: Attempt to identify and return positive prompt
|
|
- negative_only: Attempt to identify and return negative prompt
|
|
- all_prompts: Return all found prompts concatenated
|
|
|
|
Returns:
|
|
tuple: (prompt, negative_prompt, metadata_json, image_tensor)
|
|
"""
|
|
metadata = {}
|
|
warning = ""
|
|
|
|
# Get full path from the selected filename
|
|
image_path = folder_paths.get_annotated_filepath(image)
|
|
|
|
# Extract metadata from the file
|
|
if image_path and os.path.exists(image_path):
|
|
result, error = self.extract_from_path(image_path)
|
|
if error:
|
|
warning = error
|
|
elif result:
|
|
metadata = result
|
|
else:
|
|
warning = f"File not found: {image_path}"
|
|
|
|
# Ensure metadata is a dict before iterating
|
|
if metadata is None:
|
|
metadata = {}
|
|
|
|
# Load the image as tensor for output
|
|
img = Image.open(image_path)
|
|
img = img.convert("RGB")
|
|
image_tensor = torch.from_numpy(np.array(img).astype(np.float32) / 255.0).unsqueeze(0)
|
|
|
|
# Parse metadata
|
|
parsed_metadata = {}
|
|
for key, value in metadata.items():
|
|
if isinstance(value, str):
|
|
try:
|
|
parsed_metadata[key] = json.loads(value)
|
|
except json.JSONDecodeError:
|
|
parsed_metadata[key] = value
|
|
elif isinstance(value, bytes):
|
|
try:
|
|
parsed_metadata[key] = json.loads(value.decode('utf-8'))
|
|
except:
|
|
parsed_metadata[key] = value.decode('utf-8', errors='ignore')
|
|
|
|
# Find all prompts from CLIPTextEncode nodes
|
|
all_prompts = self.find_prompts_by_node_type(parsed_metadata, "CLIPTextEncode")
|
|
|
|
# Also try to find prompts using the 'text' key directly in the prompt structure
|
|
if not all_prompts:
|
|
all_prompts = self.find_strings_by_key(parsed_metadata, ["text"])
|
|
|
|
# Remove duplicates while preserving order
|
|
seen = set()
|
|
unique_prompts = []
|
|
for p in all_prompts:
|
|
if p not in seen and len(p.strip()) > 0:
|
|
seen.add(p)
|
|
unique_prompts.append(p)
|
|
all_prompts = unique_prompts
|
|
|
|
# Classify prompts
|
|
positive_prompts = []
|
|
negative_prompts = []
|
|
|
|
for prompt_text in all_prompts:
|
|
classification = self.classify_prompt(prompt_text, parsed_metadata)
|
|
if classification == "negative":
|
|
negative_prompts.append(prompt_text)
|
|
elif classification == "positive":
|
|
positive_prompts.append(prompt_text)
|
|
else:
|
|
# Unknown - assume positive if longer, negative if shorter
|
|
if len(prompt_text) > 50:
|
|
positive_prompts.append(prompt_text)
|
|
else:
|
|
negative_prompts.append(prompt_text)
|
|
|
|
# If no prompts found via node type, fall back to longest string method
|
|
if not all_prompts:
|
|
all_strings = []
|
|
for key, value in parsed_metadata.items():
|
|
all_strings.extend(self.find_all_strings(value))
|
|
|
|
# Filter out strings that are likely not prompts
|
|
# (class names, file paths, node types, etc.)
|
|
filtered_strings = []
|
|
for s in all_strings:
|
|
s_stripped = s.strip()
|
|
# Skip empty, very short, or technical strings
|
|
if len(s_stripped) < 5:
|
|
continue
|
|
# Skip strings that look like class names or paths
|
|
if s_stripped.startswith(("ComfyUI", "KSampler", "CLIP", "/", "\\", "http")):
|
|
continue
|
|
# Skip strings that are all uppercase (likely constants)
|
|
if s_stripped.isupper() and len(s_stripped) < 20:
|
|
continue
|
|
# Skip strings that look like node class names
|
|
if "_" in s_stripped and s_stripped[0].isupper() and len(s_stripped) < 30:
|
|
continue
|
|
filtered_strings.append(s_stripped)
|
|
|
|
if filtered_strings:
|
|
# Sort by length, longest first
|
|
filtered_strings.sort(key=len, reverse=True)
|
|
if len(filtered_strings) >= 1:
|
|
positive_prompts = [filtered_strings[0]]
|
|
if len(filtered_strings) >= 2:
|
|
negative_prompts = [filtered_strings[1]]
|
|
|
|
# Prepare outputs based on mode
|
|
prompt_output = ""
|
|
negative_output = ""
|
|
|
|
if extraction_mode == "longest_string":
|
|
all_found = positive_prompts + negative_prompts
|
|
if all_found:
|
|
prompt_output = max(all_found, key=len)
|
|
elif extraction_mode == "positive_only":
|
|
if positive_prompts:
|
|
prompt_output = max(positive_prompts, key=len)
|
|
elif extraction_mode == "negative_only":
|
|
if negative_prompts:
|
|
negative_output = max(negative_prompts, key=len)
|
|
elif extraction_mode == "all_prompts":
|
|
prompt_output = " | ".join(positive_prompts) if positive_prompts else ""
|
|
negative_output = " | ".join(negative_prompts) if negative_prompts else ""
|
|
|
|
# Default behavior: always try to populate both outputs
|
|
if not prompt_output and positive_prompts:
|
|
prompt_output = max(positive_prompts, key=len)
|
|
if not negative_output and negative_prompts:
|
|
negative_output = max(negative_prompts, key=len)
|
|
|
|
# Add warning to output if present
|
|
if warning and not prompt_output:
|
|
prompt_output = f"[Warning: {warning}]"
|
|
|
|
# Prepare metadata JSON output
|
|
metadata_json = json.dumps(parsed_metadata, indent=2) if parsed_metadata else "{}"
|
|
|
|
return (prompt_output, negative_output, metadata_json, image_tensor)
|
|
|
|
|
|
class DJZ_PromptExtractorFromPath:
|
|
"""
|
|
Simplified version that takes a file path directly.
|
|
Useful when you have the path to a PNG file and want to extract its prompt.
|
|
"""
|
|
|
|
def __init__(self):
|
|
pass
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"image_path": ("STRING", {
|
|
"default": "",
|
|
"multiline": False,
|
|
"placeholder": "Path to PNG file"
|
|
}),
|
|
},
|
|
"optional": {
|
|
"extraction_mode": (["longest_string", "positive_only", "negative_only", "all_prompts"],),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("STRING", "STRING", "STRING",)
|
|
RETURN_NAMES = ("prompt", "negative_prompt", "metadata_json",)
|
|
FUNCTION = "extract_prompt"
|
|
CATEGORY = "DJZ-Nodes"
|
|
OUTPUT_NODE = False
|
|
|
|
def extract_prompt(self, image_path, extraction_mode="longest_string"):
|
|
"""Extract prompt from a file path."""
|
|
metadata = {}
|
|
warning = ""
|
|
|
|
# Extract metadata from the file path
|
|
if image_path and image_path.strip():
|
|
path = image_path.strip()
|
|
if os.path.exists(path):
|
|
try:
|
|
img = Image.open(path)
|
|
metadata = img.info
|
|
except Exception as e:
|
|
warning = f"Error reading image: {str(e)}"
|
|
else:
|
|
warning = f"File not found: {path}"
|
|
|
|
# Ensure metadata is a dict
|
|
if metadata is None:
|
|
metadata = {}
|
|
|
|
# Parse metadata
|
|
parsed_metadata = {}
|
|
for key, value in metadata.items():
|
|
if isinstance(value, str):
|
|
try:
|
|
parsed_metadata[key] = json.loads(value)
|
|
except json.JSONDecodeError:
|
|
parsed_metadata[key] = value
|
|
elif isinstance(value, bytes):
|
|
try:
|
|
parsed_metadata[key] = json.loads(value.decode('utf-8'))
|
|
except:
|
|
parsed_metadata[key] = value.decode('utf-8', errors='ignore')
|
|
|
|
# Find all prompts from CLIPTextEncode nodes
|
|
all_prompts = DJZ_PromptExtractor.find_prompts_by_node_type(parsed_metadata, "CLIPTextEncode")
|
|
|
|
# Also try to find prompts using the 'text' key directly
|
|
if not all_prompts:
|
|
all_prompts = DJZ_PromptExtractor.find_strings_by_key(parsed_metadata, ["text"])
|
|
|
|
# Remove duplicates while preserving order
|
|
seen = set()
|
|
unique_prompts = []
|
|
for p in all_prompts:
|
|
if p not in seen and len(p.strip()) > 0:
|
|
seen.add(p)
|
|
unique_prompts.append(p)
|
|
all_prompts = unique_prompts
|
|
|
|
# Classify prompts
|
|
positive_prompts = []
|
|
negative_prompts = []
|
|
|
|
for prompt_text in all_prompts:
|
|
classification = DJZ_PromptExtractor.classify_prompt(prompt_text, parsed_metadata)
|
|
if classification == "negative":
|
|
negative_prompts.append(prompt_text)
|
|
elif classification == "positive":
|
|
positive_prompts.append(prompt_text)
|
|
else:
|
|
# Unknown - assume positive if longer, negative if shorter
|
|
if len(prompt_text) > 50:
|
|
positive_prompts.append(prompt_text)
|
|
else:
|
|
negative_prompts.append(prompt_text)
|
|
|
|
# If no prompts found via node type, fall back to longest string method
|
|
if not all_prompts:
|
|
all_strings = []
|
|
for key, value in parsed_metadata.items():
|
|
all_strings.extend(DJZ_PromptExtractor.find_all_strings(value))
|
|
|
|
# Filter out strings that are likely not prompts
|
|
filtered_strings = []
|
|
for s in all_strings:
|
|
s_stripped = s.strip()
|
|
if len(s_stripped) < 5:
|
|
continue
|
|
if s_stripped.startswith(("ComfyUI", "KSampler", "CLIP", "/", "\\", "http")):
|
|
continue
|
|
if s_stripped.isupper() and len(s_stripped) < 20:
|
|
continue
|
|
if "_" in s_stripped and s_stripped[0].isupper() and len(s_stripped) < 30:
|
|
continue
|
|
filtered_strings.append(s_stripped)
|
|
|
|
if filtered_strings:
|
|
filtered_strings.sort(key=len, reverse=True)
|
|
if len(filtered_strings) >= 1:
|
|
positive_prompts = [filtered_strings[0]]
|
|
if len(filtered_strings) >= 2:
|
|
negative_prompts = [filtered_strings[1]]
|
|
|
|
# Prepare outputs based on mode
|
|
prompt_output = ""
|
|
negative_output = ""
|
|
|
|
if extraction_mode == "longest_string":
|
|
all_found = positive_prompts + negative_prompts
|
|
if all_found:
|
|
prompt_output = max(all_found, key=len)
|
|
elif extraction_mode == "positive_only":
|
|
if positive_prompts:
|
|
prompt_output = max(positive_prompts, key=len)
|
|
elif extraction_mode == "negative_only":
|
|
if negative_prompts:
|
|
negative_output = max(negative_prompts, key=len)
|
|
elif extraction_mode == "all_prompts":
|
|
prompt_output = " | ".join(positive_prompts) if positive_prompts else ""
|
|
negative_output = " | ".join(negative_prompts) if negative_prompts else ""
|
|
|
|
# Default behavior: always try to populate both outputs
|
|
if not prompt_output and positive_prompts:
|
|
prompt_output = max(positive_prompts, key=len)
|
|
if not negative_output and negative_prompts:
|
|
negative_output = max(negative_prompts, key=len)
|
|
|
|
# Add warning to output if present
|
|
if warning and not prompt_output:
|
|
prompt_output = f"[Warning: {warning}]"
|
|
|
|
# Prepare metadata JSON output
|
|
metadata_json = json.dumps(parsed_metadata, indent=2) if parsed_metadata else "{}"
|
|
|
|
return (prompt_output, negative_output, metadata_json)
|
|
|
|
|
|
class DJZ_PromptExtractorBatch:
|
|
"""
|
|
Batch version that processes multiple images from a directory.
|
|
Returns prompts as a list/batch.
|
|
"""
|
|
|
|
def __init__(self):
|
|
pass
|
|
|
|
@classmethod
|
|
def INPUT_TYPES(cls):
|
|
return {
|
|
"required": {
|
|
"directory_path": ("STRING", {
|
|
"default": "",
|
|
"multiline": False,
|
|
"placeholder": "Path to directory containing PNG files"
|
|
}),
|
|
},
|
|
"optional": {
|
|
"max_images": ("INT", {
|
|
"default": 10,
|
|
"min": 1,
|
|
"max": 100,
|
|
"step": 1
|
|
}),
|
|
"extraction_mode": (["longest_string", "positive_only", "negative_only"],),
|
|
}
|
|
}
|
|
|
|
RETURN_TYPES = ("STRING", "STRING",)
|
|
RETURN_NAMES = ("prompts_list", "filenames_list",)
|
|
FUNCTION = "extract_batch"
|
|
CATEGORY = "DJZ-Nodes"
|
|
OUTPUT_NODE = False
|
|
|
|
def extract_batch(self, directory_path, max_images=10, extraction_mode="longest_string"):
|
|
"""Extract prompts from all PNG files in a directory."""
|
|
if not os.path.isdir(directory_path):
|
|
return (f"[Error: Directory not found: {directory_path}]", "")
|
|
|
|
# Find all PNG files
|
|
png_files = [f for f in os.listdir(directory_path) if f.lower().endswith('.png')]
|
|
png_files = png_files[:max_images]
|
|
|
|
if not png_files:
|
|
return ("[No PNG files found in directory]", "")
|
|
|
|
prompts = []
|
|
filenames = []
|
|
path_extractor = DJZ_PromptExtractorFromPath()
|
|
|
|
for filename in png_files:
|
|
filepath = os.path.join(directory_path, filename)
|
|
prompt, _, _ = path_extractor.extract_prompt(filepath, extraction_mode)
|
|
if prompt and not prompt.startswith("["):
|
|
prompts.append(prompt)
|
|
filenames.append(filename)
|
|
|
|
# Join with newlines for easy reading
|
|
prompts_output = "\n---\n".join(prompts) if prompts else "[No prompts extracted]"
|
|
filenames_output = "\n".join(filenames) if filenames else ""
|
|
|
|
return (prompts_output, filenames_output)
|
|
|
|
|
|
# Node registration for ComfyUI
|
|
NODE_CLASS_MAPPINGS = {
|
|
"DJZ_PromptExtractor": DJZ_PromptExtractor,
|
|
"DJZ_PromptExtractorFromPath": DJZ_PromptExtractorFromPath,
|
|
"DJZ_PromptExtractorBatch": DJZ_PromptExtractorBatch,
|
|
}
|
|
|
|
NODE_DISPLAY_NAME_MAPPINGS = {
|
|
"DJZ_PromptExtractor": "DJZ Prompt Extractor",
|
|
"DJZ_PromptExtractorFromPath": "DJZ Prompt Extractor (Path)",
|
|
"DJZ_PromptExtractorBatch": "DJZ Prompt Extractor (Batch)",
|
|
}
|