NodeAware Update

This commit is contained in:
holonic
2024-06-02 01:06:56 +01:00
parent 49060b0b4f
commit 3d4900e65d
4 changed files with 486 additions and 270 deletions
+1
View File
@@ -1,3 +1,4 @@
__pycache__
test
notes
.vscode
+383 -269
View File
@@ -11,8 +11,8 @@
# - How do I use just the "internal variables" since it seems scope on AnyNode class params are global (not good for storing script)
# - Security on globals
# - use re for parsing the python code instead of naively expecting the response to start with a python tag
# - fix gemini node, give it some love
# Import common libs
import os, json, random, string, sys, math, datetime, collections, itertools, functools, urllib, shutil, re, torch, time, decimal, matplotlib, io, base64, wave, chromadb, uuid
import numpy
import numpy as np
@@ -27,12 +27,31 @@ import google.generativeai as genai
import pkgutil
import importlib
from openai import OpenAI
# from .context_utils import is_context_empty, _create_context_data
# from .constants import get_category, get_name
# Comfy libs
def add_comfy_path():
current_path = os.path.dirname(os.path.abspath(__file__))
comfy_path = os.path.abspath(os.path.join(current_path, '../../../comfy'))
if comfy_path not in sys.path:
sys.path.insert(0, comfy_path)
add_comfy_path()
import comfy.diffusers_load # type: ignore
import comfy.samplers # type: ignore
import comfy.sample # type: ignore
import comfy.sd # type: ignore
import comfy.utils # type: ignore
import comfy.controlnet # type: ignore
import comfy.clip_vision # type: ignore
import comfy.model_management # type: ignore
# Packaged Utility libs
from .utils import any_type, is_none, variable_info, sanitize_code
from .util_gemini import GoogleGemini
from .util_oai_compatible import OpenAICompatible
from .util_functions import FunctionRegistry
from .util_nodeaware import NodeAware
# The template for the system message sent to ChatCompletions
SYSTEM_TEMPLATE = """
@@ -47,7 +66,7 @@ It is not required to use any of these libraries, but if you do use any import i
Here is some important information about the input data:
- input_data_1: [[INPUT1]]
- input_data_2: [[INPUT2]]
[[EXAMPLES]][[CODEBLOCK]]
[[CONNECTIONS]][[EXAMPLES]][[CODEBLOCK]]
## Coding Instructions
- Your job is to code the user's requested node given the inputs and desired output type.
- Respond with only a brief plan and the code in one function named generated_function that takes two kwargs named 'input_data_1' and 'input_data_2'.
@@ -90,284 +109,305 @@ def generated_function(input_data_1=None, input_data_2=None):
"""
class AnyNode:
"""Ask it to make up any node for you. """
NAME = "AnyNode"
CATEGORY = "utils"
ALLOWED_IMPORTS = {"os", "re", "json", "random", "string", "sys", "math", "datetime", "collections", "itertools", "functools", "numpy", "openai", "traceback", "torch", "time", "sklearn", "torchvision", "matplotlib", "io", "base64", "wave", "google.generativeai", "chromadb", "uuid"}
CODING_ATTEMPTS = 3
def __init__(self):
self.oai_model = "gpt-4o"
self.reset()
self.unique_id = str(uuid.uuid4()).replace('-', '')
def generate_function_name(self):
return f"generated_function_{self.unique_id}"
def reset(self):
self.script = None
self.last_prompt = None
self.imports:list[str] = []
self.last_error = None
self.last_comment = None
self.attempts = 0
"""Ask it to make up any node for you. """
@classmethod
def INPUT_TYPES(self): # pylint: disable = invalid-name, missing-function-docstring
return {
"required": {
"prompt": ("STRING", {
"multiline": True,
"default": "Take the input and multiply by 5",
}),
"model": (["gpt-4o", "gpt-4-turbo", "gpt-4", "gpt-3.5-turbo", "gpt-3.5"], {
"default": "gpt-4o"
}),
},
"optional": {
"any": (any_type,),
"any2": (any_type,),
},
}
@classmethod
def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen):
return time.time()
RETURN_TYPES = (any_type,)
RETURN_NAMES = ('any',)
FUNCTION = "go"
OUTPUT_NODE = True
NAME = "AnyNode"
CATEGORY = "utils"
# TODO: Store the md5 of a prompt in function cache globally so that a duplicated node will not need to resolve
# store function cache JSON in 'output' folder!!!!! baller.
VERSION = "0.1.2"
FUNCTION_REGISTRY = FunctionRegistry(schema="default", version=VERSION)
ALLOWED_IMPORTS = {"os", "re", "json", "random", "string", "sys", "math", "datetime", "collections", "itertools", "functools", "numpy", "openai", "traceback", "torch", "time", "sklearn", "torchvision", "matplotlib", "io", "base64", "wave", "google.generativeai", "chromadb", "uuid", "comfy"}
CODING_ATTEMPTS = 3
def render_template(self, template:str, any=None, any2=None, seed=None):
"""Render the system template with current state"""
varinfo = [variable_info(any), variable_info(any2)]
print(f"LE: {self.last_error}")
instruction = "" if not self.last_error else f"There was an error with the last generated_function.\n\n### Debugging Instructions\n-If the error is that something is 'not defined' find a workaround using an alternative.\n- If the undefined thing is a function, most likely you didn't wrap the function inside `generated_function`.\n- Reflect on the error in your reply, be concise and accurate in analyzing the problem, then write the updated generated_function.\n- If there is a ValueError about a Dangerous construct being detected, your code has not passed the sanitizer; find an alternative.\n\n### Traceback\n{self.last_error}\n\n### Erroneous Code"
#print(f"Input 0 -> {varinfo}")
examples = ""
r = template \
.replace('[[IMPORTS]]', ", ".join(list(self.ALLOWED_IMPORTS))) \
.replace('[[INPUT1]]', varinfo[0]) \
.replace('[[INPUT2]]', varinfo[1]) \
.replace('[[EXAMPLES]]', examples) \
.replace("[[CODEBLOCK]]", "" if not self.script else f"\n## Current Code\n{instruction}\n```python\n{self.script}\n```\n")
# This is the case where we call from error mitigation
return r
def get_response(self, system_message:str, prompt:str, model=None, **kwargs) -> str:
"""Calls OpenAI With System Message and Prompt. Overriden in classes that extend this."""
if model:
self.oai_model = model
client = OpenAI(
# This is the default and can be omitted
api_key=os.environ.get("OPENAI_API_KEY"),
)
response = client.chat.completions.create(
model=self.oai_model, # Use the model of your choice, e.g., gpt-4 or gpt-3.5-turbo
messages=[
{"role": "system", "content": system_message},
{"role": "user", "content": prompt}
]
)
# Extract the response text
r = response.choices[0].message.content.strip()
return r
def get_llm_response(self, prompt:str, any=None, any2=None, **kwargs) -> str:
"""Calls OpenAI and returns response"""
try:
print(f"INPUT ({type(any)}, {type(any2)})")
final_template = self.render_template(SYSTEM_TEMPLATE, any=any, any2=any2)
print(final_template)
r = self.get_response(final_template, prompt, **kwargs)
code_block = self.extract_code_block(r)
print(f"LLM COMMENTS:\n{self.last_comment}")
return code_block
except Exception as e:
return f"An error occurred: {e}"
def extract_code_block(self, response: str) -> str:
"""
Extracts the code block from the response using regex.
Saves everything before the code block as self.last_comment.
Returns the code block as a string.
"""
code_pattern = re.compile(r'(.*?)```python(.*?)```', re.DOTALL)
match = code_pattern.search(response)
if match:
self.last_comment = match.group(1).strip()
return match.group(2).strip()
else:
self.last_comment = response.strip('`')
return self.last_comment
def extract_imports(self, generated_code):
"""
Extracts import statements from the generated code and stores them in self.imports.
Returns the code without the import statements.
"""
import_pattern = re.compile(r'^\s*(import .+|from .+ import .+)', re.MULTILINE)
imports = import_pattern.findall(generated_code)
cleaned_code = import_pattern.sub('', generated_code).strip()
# Store the imports in the instance variable
self.imports = [imp.strip() for imp in imports]
print(f"Imports in code: {self.imports}")
return cleaned_code
def _prepare_globals(self, globals_dict: dict):
"""Get the globals dict prepared for safe_exec"""
for imp in self.imports:
parts = imp.split()
try:
if imp.startswith('import'):
# Handle 'import module'
if len(parts) == 2:
module_name = parts[1]
globals_dict[module_name] = importlib.import_module(module_name)
self.import_submodules(module_name, globals_dict)
# Handle 'import module as alias'
elif len(parts) == 4 and parts[2] == 'as':
module_name = parts[1]
alias = parts[3]
globals_dict[alias] = importlib.import_module(module_name)
self.import_submodules(module_name, globals_dict)
elif imp.startswith('from'):
# Handle 'from module import name'
if len(parts) == 4:
module_name = parts[1]
name = parts[3]
globals_dict[name] = importlib.import_module(f"{module_name}.{name}")
# Handle 'from module import name as alias'
elif len(parts) == 6 and parts[4] == 'as':
module_name = parts[1]
name = parts[3]
alias = parts[5]
globals_dict[alias] = importlib.import_module(f"{module_name}.{name}")
except ImportError as e:
print(f"Failed to import {imp}: {e}")
def import_submodules(self, package_name, globals_dict):
"""Get the submodules from a package and import those into the globals"""
if package_name in sys.modules:
package = sys.modules[package_name]
if hasattr(package, '__path__'):
for loader, module_name, is_pkg in pkgutil.walk_packages(package.__path__, package.__name__ + '.'):
if any(submodule.startswith(module_name) for submodule in self.imports):
try:
module = importlib.import_module(module_name)
globals_dict[module_name] = module
except ImportError as e:
print(f"Failed to import submodule {module_name}: {e}")
traceback.print_exc()
def safe_exec(self, code_string, globals_dict=None, locals_dict=None):
"""Execute """
if globals_dict is None:
globals_dict = {}
if locals_dict is None:
locals_dict = {}
# Import submodules for each module in globals_dict
for module_name in list(globals_dict.keys()):
try:
self.import_submodules(module_name, globals_dict)
except Exception as e:
print(f"Failed to import submodules for {module_name}: {e}")
try:
exec(sanitize_code(code_string), globals_dict, locals_dict)
except Exception as e:
print("An error occurred:")
traceback.print_exc()
raise e
def keep_trying(self):
return self.attempts < self.CODING_ATTEMPTS
def go(self, prompt:str, any=None, any2=None, **kwargs):
print("TESTTEST", prompt, any, any2)
"""Takes the prompt and inputs, Generates a function with an LLM for the Node"""
if prompt == "": # if empty, reset
def __init__(self):
self.oai_model = "gpt-4o"
self.reset()
return (any, any2,)
result = None
registry = self.FUNCTION_REGISTRY
# Generate a unique function name
function_name = self.generate_function_name()
self.unique_id = str(uuid.uuid4()).replace('-', '')
def generate_function_name(self):
return f"generated_function_{self.unique_id}"
def reset(self):
self.script = None
self.last_prompt = None
self.imports:list[str] = []
self.last_error = None
self.last_comment = None
self.attempts = 0
@classmethod
def INPUT_TYPES(self): # pylint: disable = invalid-name, missing-function-docstring
return {
"required": {
"prompt": ("STRING", {
"multiline": True,
"default": "Take the input and multiply by 5",
}),
"model": (["gpt-4o", "gpt-4-turbo", "gpt-4", "gpt-3.5-turbo", "gpt-3.5"], {
"default": "gpt-4o"
}),
},
"optional": {
"any": (any_type,),
"any2": (any_type,),
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO",
},
}
@classmethod
def IS_CHANGED(s, image, string_field, int_field, float_field, print_to_screen):
return time.time()
# Generate, Compile and Run the Unique Generated Function: 3 Attempts
while self.keep_trying():
print(f"Last Error: {self.last_error}")
fr = registry.get_function(prompt)
use_function = fr is not None and self.last_error is None
use_generation = self.script is None or self.last_prompt != prompt or self.last_error is not None
if use_generation and not use_function:
print("Generating Node function...")
# Generate the function code using OpenAI
r = self.get_llm_response(prompt, any=any, any2=any2, **kwargs)
# Remember the script for future use
self.script = self.extract_imports(r)
print(f"Stored script:\n{self.script}")
if use_function:
self.script = fr['function']
self.last_comment = fr['comment']
self.imports = fr['imports']
self.last_prompt = prompt
# Modify the script to use the unique function name
modified_script = self.script.replace('def generated_function', f'def {function_name}')
# Execute the stored script to define the unique function
RETURN_TYPES = (any_type, 'CTRL',)
RETURN_NAMES = ('any', 'control',)
FUNCTION = "go"
OUTPUT_NODE = True
# TODO: Store the md5 of a prompt in function cache globally so that a duplicated node will not need to resolve
# store function cache JSON in 'output' folder!!!!! baller.
VERSION = "0.1.2"
FUNCTION_REGISTRY = FunctionRegistry(schema="default", version=VERSION)
def render_template(self, template:str, any=None, any2=None, seed=None, workflow:NodeAware=None, node=None):
"""Render the system prompt template with current state"""
varinfo = [variable_info(any), variable_info(any2)]
print(f"LE: {self.last_error}")
instruction = "" if not self.last_error else f"There was an error with the last generated_function.\n\n### Debugging Instructions\n-If the error is that something is 'not defined' find a workaround using an alternative.\n- If the undefined thing is a function, most likely you didn't wrap the function inside `generated_function`.\n- Reflect on the error in your reply, be concise and accurate in analyzing the problem, then write the updated generated_function.\n- If there is a ValueError about a Dangerous construct being detected, your code has not passed the sanitizer; find an alternative.\n\n### Traceback\n{self.last_error}\n\n### Erroneous Code"
#print(f"Input 0 -> {varinfo}")
summary = None if not node or not workflow else workflow.summarize_connections(node['id'])
examples = ""
r = template \
.replace('[[IMPORTS]]', ", ".join(list(self.ALLOWED_IMPORTS))) \
.replace('[[INPUT1]]', varinfo[0]) \
.replace('[[INPUT2]]', varinfo[1]) \
.replace('[[CONNECTIONS]]', summary) \
.replace('[[EXAMPLES]]', examples) \
.replace("[[CODEBLOCK]]", "" if not self.script else f"\n## Current Code\n{instruction}\n```python\n{self.script}\n```\n")
# This is the case where we call from error mitigation
return r
def get_response(self, system_message:str, prompt:str, model=None, **kwargs) -> str:
"""Calls OpenAI With System Message and Prompt. Overriden in classes that extend this."""
if model:
self.oai_model = model
client = OpenAI(
# This is the default and can be omitted
api_key=os.environ.get("OPENAI_API_KEY"),
)
response = client.chat.completions.create(
model=self.oai_model, # Use the model of your choice, e.g., gpt-4 or gpt-3.5-turbo
messages=[
{"role": "system", "content": system_message},
{"role": "user", "content": prompt}
]
)
# Extract the response text
r = response.choices[0].message.content.strip()
return r
def get_llm_response(self, prompt:str, any=None, any2=None, workflow=None, node=None, **kwargs) -> str:
"""Calls OpenAI and returns response"""
try:
# Define a dictionary to store globals and locals, updating it with imported libs from script and built in functions
globals_dict = {"__builtins__": __builtins__}
self._prepare_globals(globals_dict)
globals_dict.update({"np": np})
locals_dict = {}
self.safe_exec(modified_script, globals_dict, locals_dict)
print(f"INPUT ({type(any)}, {type(any2)})")
final_template = self.render_template(SYSTEM_TEMPLATE, any=any, any2=any2, workflow=workflow, node=node)
print(final_template)
r = self.get_response(final_template, prompt, **kwargs)
code_block = self.extract_code_block(r)
print(f"LLM COMMENTS:\n{self.last_comment}")
return code_block
except Exception as e:
print("--- Exception During Exec ---")
# store the error for next run
self.last_error = traceback.format_exc()
if not self.keep_trying():
raise e
continue
# Assuming the generated code defines a function named 'generated_function'
if function_name in locals_dict:
return f"An error occurred: {e}"
def extract_code_block(self, response: str) -> str:
"""
Extracts the code block from the response using regex.
Saves everything before the code block as self.last_comment.
Returns the code block as a string.
"""
code_pattern = re.compile(r'(.*?)```python(.*?)```', re.DOTALL)
match = code_pattern.search(response)
if match:
self.last_comment = match.group(1).strip()
return match.group(2).strip()
else:
self.last_comment = response.strip('`')
return self.last_comment
def extract_imports(self, generated_code):
"""
Extracts import statements from the generated code and stores them in self.imports.
Returns the code without the import statements.
"""
import_pattern = re.compile(r'^\s*(import .+|from .+ import .+)', re.MULTILINE)
imports = import_pattern.findall(generated_code)
cleaned_code = import_pattern.sub('', generated_code).strip()
# Store the imports in the instance variable
self.imports = [imp.strip() for imp in imports]
print(f"Imports in code: {self.imports}")
return cleaned_code
def _prepare_globals(self, globals_dict: dict):
"""Get the globals dict prepared for safe_exec"""
for imp in self.imports:
parts = imp.split()
try:
# Call the generated function and get the result
result = locals_dict[function_name](any, input_data_2=any2)
print(f"Function result: {result}")
if imp.startswith('import'):
# Handle 'import module'
if len(parts) == 2:
module_name = parts[1]
globals_dict[module_name] = importlib.import_module(module_name)
self.import_submodules(module_name, globals_dict)
# Handle 'import module as alias'
elif len(parts) == 4 and parts[2] == 'as':
module_name = parts[1]
alias = parts[3]
globals_dict[alias] = importlib.import_module(module_name)
self.import_submodules(module_name, globals_dict)
elif imp.startswith('from'):
# Handle 'from module import name'
if len(parts) == 4:
module_name = parts[1]
name = parts[3]
globals_dict[name] = importlib.import_module(f"{module_name}.{name}")
# Handle 'from module import name as alias'
elif len(parts) == 6 and parts[4] == 'as':
module_name = parts[1]
name = parts[3]
alias = parts[5]
globals_dict[alias] = importlib.import_module(f"{module_name}.{name}")
except ImportError as e:
print(f"Failed to import {imp}: {e}")
def import_submodules(self, package_name, globals_dict):
"""Get the submodules from a package and import those into the globals"""
if package_name in sys.modules:
package = sys.modules[package_name]
if hasattr(package, '__path__'):
for loader, module_name, is_pkg in pkgutil.walk_packages(package.__path__, package.__name__ + '.'):
if any(submodule.startswith(module_name) for submodule in self.imports):
try:
module = importlib.import_module(module_name)
globals_dict[module_name] = module
except ImportError as e:
print(f"Failed to import submodule {module_name}: {e}")
traceback.print_exc()
def safe_exec(self, code_string, globals_dict=None, locals_dict=None):
"""Execute """
if globals_dict is None:
globals_dict = {}
if locals_dict is None:
locals_dict = {}
# Import submodules for each module in globals_dict
for module_name in list(globals_dict.keys()):
try:
self.import_submodules(module_name, globals_dict)
except Exception as e:
print(f"Error calling the generated function: {e}")
traceback.print_exc()
print(f"Failed to import submodules for {module_name}: {e}")
try:
exec(sanitize_code(code_string), globals_dict, locals_dict)
except Exception as e:
print("An error occurred:")
traceback.print_exc()
raise e
def keep_trying(self):
r = self.attempts < self.CODING_ATTEMPTS
self.attempts += 1
return r
def go(self, prompt:str, any=None, any2=None, unique_id=None, extra_pnginfo=None, **kwargs):
print("TESTTEST", prompt, any, any2)
"""Takes the prompt and inputs, Generates a function with an LLM for the Node"""
if prompt == "": # if empty, reset
self.reset()
return (any, any2,)
result = None
registry = self.FUNCTION_REGISTRY
# Generate a unique function name
function_name = self.generate_function_name()
workflow = NodeAware(pnginfo=extra_pnginfo)
node = workflow.find_node(id=unique_id)
# Generate, Compile and Run the Unique Generated Function: 3 Attempts
while self.keep_trying():
print(f"Last Error: {self.last_error}")
fr = registry.get_function(prompt)
use_function = fr is not None and self.last_error is None
use_generation = self.script is None or self.last_prompt != prompt or self.last_error is not None
if use_generation and not use_function:
print("Generating Node function...")
# Generate the function code using OpenAI
r = self.get_llm_response(prompt, any=any, any2=any2, workflow=workflow, node=node, **kwargs)
# Remember the script for future use
self.script = self.extract_imports(r)
print(f"Stored script:\n{self.script}")
if use_function:
self.script = fr['function']
self.last_comment = fr['comment']
self.imports = fr['imports']
self.last_prompt = prompt
# Modify the script to use the unique function name
modified_script = self.script.replace('def generated_function', f'def {function_name}')
# Execute the stored script to define the unique function
try:
# Define a dictionary to store globals and locals, updating it with imported libs from script and built in functions
globals_dict = {"__builtins__": __builtins__}
self._prepare_globals(globals_dict)
globals_dict.update({"np": np})
locals_dict = {}
self.safe_exec(modified_script, globals_dict, locals_dict)
except Exception as e:
print("--- Exception During Exec ---")
# store the error for next run
self.last_error = traceback.format_exc()
if not self.keep_trying():
raise e
continue
else:
print(f"Function '{function_name}' not found in generated code.")
break
self.last_error = None
# Here we assume the function is complete and we can store it in the registry
registry.add_function(prompt, self.script, self.imports, self.last_comment, [variable_info(any), variable_info(any2)])
self.attempts = 0
return (result,)
# Assuming the generated code defines a function named 'generated_function'
if function_name in locals_dict:
try:
# Call the generated function and get the result
result = locals_dict[function_name](any, input_data_2=any2)
print(f"Function result: {result}")
except Exception as e:
print(f"Error calling the generated function: {e}")
traceback.print_exc()
self.last_error = traceback.format_exc()
if not self.keep_trying():
raise e
continue
else:
print(f"Function '{function_name}' not found in generated code.")
break
self.last_error = None
# Here we assume the function is complete and we can store it in the registry
registry.add_function(prompt, self.script, self.imports, self.last_comment, [variable_info(any), variable_info(any2)])
self.attempts = 0
# Control data
control = {
'prompt': self.last_prompt,
'last_comment': self.last_comment,
'inputs': (variable_info(any), variable_info(any2)),
'imports': self.imports,
'function': self.script,
'last_error': self.last_error,
}
return (result, control,)
class AnyNodeGemini(AnyNode):
def __init__(self, api_key=None):
super().__init__()
@@ -430,19 +470,93 @@ class AnyNodeOpenAICompatible(AnyNode):
self.llm.set_api_server(server)
return self.llm.get_response(system_message, prompt, any=any)
class AnyNodeCodeViewer:
@classmethod
def INPUT_TYPES(s):
return {
"required": {
"ctrl": ("CTRL", {"forceInput": True}),
},
"optional": {
"text": ("STRING", {
"multiline": True,
"default": "No Comment yet.",
}),
},
"hidden": {
"unique_id": "UNIQUE_ID",
"extra_pnginfo": "EXTRA_PNGINFO",
},
}
#INPUT_IS_LIST = True
RETURN_TYPES = ("CTRL",)
RETURN_NAMES = ("control",)
FUNCTION = "notify"
OUTPUT_NODE = True
#OUTPUT_IS_LIST = (True,)
CATEGORY = "utils"
def notify(self, ctrl, text=None, unique_id=None, extra_pnginfo=None):
# Ensure ctrl is a dictionary
if isinstance(unique_id, list):
unique_id = unique_id[0]
if not isinstance(ctrl, dict):
return {"ui": {"text": "Error: Input is not a dictionary"}, "result": {}}
# Extract information from the ctrl dictionary
prompt = ctrl.get("prompt", "No prompt provided")
last_comment = ctrl.get("last_comment", "No comment provided")
inputs = ctrl.get("inputs", "No inputs provided")
imports = ctrl.get("imports", "No imports provided")
function = ctrl.get("function", "No function provided")
last_error = ctrl.get("last_error", "No error provided")
# Prepare the display text
display_text = f"## LLM Output\n\n### Last Comment\n{last_comment}\n\n"
display_text += f"### Function\n```python\n{function}\n```\n\n"
display_text += f"### Imports\n{', '.join(imports)}\n\n"
display_text += f"### Last Error\n{last_error}\n\n"
# Update the workflow with ctrl information
has_uid = unique_id is not None
has_png = extra_pnginfo is not None
# Directly manipulate the workflow to show text on this node
if has_uid and has_png:
# Find the node
node_aware = NodeAware(pnginfo=extra_pnginfo)
print("\nWORKFLOW", node_aware.workflow)
node = node_aware.find_node(id=unique_id)
print("\nNODE", node, '\n')
# Show the display text by adding a widget value
if node:
node["widgets_values"] = [display_text]
print("AFTER", extra_pnginfo['workflow']['nodes'])
#return (ctrl,)
return {"ui": {"text": display_text}, "result": (ctrl,)}
#return {"ui": {"text": display_text}, "result": ctrl}
# Manager Mappings
NODE_CLASS_MAPPINGS = {
"AnyNode": AnyNode,
"AnyNodeGemini": AnyNodeGemini,
"AnyNodeLocal": AnyNodeOpenAICompatible,
#"AnyNodeCodeViewer": AnyNodeCodeViewer,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"AnyNode": "Any Node 🍄",
"AnyNodeGemini": "Any Node 🍄 (Gemini)",
"AnyNodeLocal": "Any Node 🍄 (Local LLM)",
#"AnyNodeCodeViewer": "View Code 🍄 - AnyNode"
}
# Unit test
if __name__ == "__main__":
node = AnyNode()
example_prompt = "Generate a random number using the input as seed"
+101
View File
@@ -0,0 +1,101 @@
class NodeAware:
"""Utility class for working with ComfyUI workflow data from hidden `prompt` in your custom nodes inputs"""
def __init__(self, workflow=None, pnginfo=None):
if pnginfo:
if not isinstance(pnginfo, list):
if isinstance(pnginfo, dict):
workflow = pnginfo["workflow"]
else:
workflow = pnginfo[0]["workflow"]
self.workflow = workflow
def get_node(self, node_id):
return next((node for node in self.workflow['nodes'] if node['id'] == node_id), None)
def get_connected_nodes(self, node_id, direction='both'):
"""Find all connected nodes and what slots they are connected to"""
connections = {'inputs': [], 'outputs': []}
for link in self.workflow['links']:
if direction in ('both', 'inputs') and link[3] == node_id:
source_node = self.get_node(link[1])
connections['inputs'].append({
'source_node': link[1],
'source_node_name': source_node['type'],
'source_slot': link[2],
'target_node': link[3],
'target_slot': link[4],
'type': link[5]
})
if direction in ('both', 'outputs') and link[1] == node_id:
target_node = self.get_node(link[3])
connections['outputs'].append({
'target_node': link[3],
'target_node_name': target_node['type'],
'target_slot': link[4],
'source_node': link[1],
'source_slot': link[2],
'type': link[5]
})
return connections
def summarize_connections(self, node_id):
"""Summarize the connections so that an LLM can have situational context for it's system prompt"""
connections = self.get_connected_nodes(node_id)
summary = f"### Node {node_id} Connections\n\n"
summary += "#### Inputs\n"
for input_conn in connections['inputs']:
summary += f"- From Node {input_conn['source_node']} ({input_conn['source_node_name']}) on slot {input_conn['source_slot']} (type: {input_conn['type']})\n"
summary += "\n#### Outputs\n"
for output_conn in connections['outputs']:
summary += f"- To Node {output_conn['target_node']} ({output_conn['target_node_name']}) on slot {output_conn['target_slot']} (type: {output_conn['type']})\n"
return summary
def find_nodes(self, **kwargs):
"""Search for nodes by id, type, input_types:list, or output_types:list. Returns a *list of the node objects from the workflow"""
print(f"\nFinding Nodes in Workspace. kwargs: {kwargs}")
result = []
node_id = kwargs.get('id')
node_type = kwargs.get('type')
input_types = kwargs.get('input_types', [])
output_types = kwargs.get('output_types', [])
for node in self.workflow['nodes']:
if node_id and int(node['id']) == int(node_id):
result.append(node)
elif node_type and node['type'] == node_type:
result.append(node)
elif input_types:
for input_link in node['inputs']:
if input_link['type'] in input_types:
result.append(node)
break
elif output_types:
for output_link in node['outputs']:
if output_link['type'] in output_types:
result.append(node)
break
return None if len(result) == 0 else result
def find_node(self, **kwargs):
r = self.find_nodes(**kwargs)
if r:
return r[0]
return r
if __name__ == "__main__":
# Example usage
workflow_data = {'last_node_id': 13, 'last_link_id': 13, 'nodes': [{'id': 5, 'type': 'ShowText|pysssss', 'pos': [1036.8042600016947, 84.57110116501482], 'size': [305.0083925861959, 219.76858999709455], 'flags': {}, 'order': 7, 'mode': 0, 'inputs': [{'name': 'text', 'type': 'STRING', 'link': 1, 'widget': {'name': 'text'}}], 'outputs': [{'name': 'STRING', 'type': 'STRING', 'links': None, 'shape': 6}], 'properties': {'Node name for S&R': 'ShowText|pysssss'}, 'widgets_values': ['', 'Input 1 -> Type: dict, Keys: samples\nInput 2 -> Type: Tensor']}, {'id': 3, 'type': 'LoadVideo [n-suite]', 'pos': [214, 64], 'size': [210, 630], 'flags': {}, 'order': 0, 'mode': 0, 'outputs': [{'name': 'IMAGES', 'type': 'IMAGE', 'links': [5], 'shape': 6, 'slot_index': 0}, {'name': 'EMPTY LATENTS', 'type': 'LATENT', 'links': None, 'shape': 6}, {'name': 'METADATA', 'type': 'STRING', 'links': [], 'shape': 3, 'slot_index': 2}, {'name': 'WIDTH', 'type': 'INT', 'links': [], 'shape': 3, 'slot_index': 3}, {'name': 'HEIGHT', 'type': 'INT', 'links': None, 'shape': 3}, {'name': 'META_FPS', 'type': 'INT', 'links': [], 'shape': 3, 'slot_index': 5}, {'name': 'META_N_FRAMES', 'type': 'INT', 'links': None, 'shape': 3}], 'properties': {'Node name for S&R': 'LoadVideo [n-suite]'}, 'widgets_values': ['swan-lake-tchaikovski.mp4', '/view?filename=swan-lake-tchaikovski.mp4&type=input&subfolder=n-suite', 'original', 'none', 512, 0, 0, 0, True, 'image', None]}, {'id': 8, 'type': 'CheckpointLoaderSimple', 'pos': [-333, -268], 'size': {'0': 315, '1': 98}, 'flags': {}, 'order': 1, 'mode': 0, 'outputs': [{'name': 'MODEL', 'type': 'MODEL', 'links': [7], 'shape': 3}, {'name': 'CLIP', 'type': 'CLIP', 'links': [8, 9], 'shape': 3, 'slot_index': 1}, {'name': 'VAE', 'type': 'VAE', 'links': None, 'shape': 3}], 'properties': {'Node name for S&R': 'CheckpointLoaderSimple'}, 'widgets_values': ['3dMixCharacter_v20Realism.safetensors']}, {'id': 10, 'type': 'CLIPTextEncode', 'pos': [-149, 17], 'size': [250, 88], 'flags': {}, 'order': 4, 'mode': 0, 'inputs': [{'name': 'clip', 'type': 'CLIP', 'link': 9}], 'outputs': [{'name': 'CONDITIONING', 'type': 'CONDITIONING', 'links': [11], 'shape': 3, 'slot_index': 0}], 'properties': {'Node name for S&R': 'CLIPTextEncode'}, 'widgets_values': ['']}, {'id': 9, 'type': 'CLIPTextEncode', 'pos': [-142, -104], 'size': [234, 78], 'flags': {}, 'order': 3, 'mode': 0, 'inputs': [{'name': 'clip', 'type': 'CLIP', 'link': 8}], 'outputs': [{'name': 'CONDITIONING', 'type': 'CONDITIONING', 'links': [10], 'shape': 3, 'slot_index': 0}], 'properties': {'Node name for S&R': 'CLIPTextEncode'}, 'widgets_values': ['a cat']}, {'id': 7, 'type': 'KSampler', 'pos': [299, -230], 'size': {'0': 315, '1': 262}, 'flags': {}, 'order': 5, 'mode': 0, 'inputs': [{'name': 'model', 'type': 'MODEL', 'link': 7, 'slot_index': 0}, {'name': 'positive', 'type': 'CONDITIONING', 'link': 10}, {'name': 'negative', 'type': 'CONDITIONING', 'link': 11}, {'name': 'latent_image', 'type': 'LATENT', 'link': 12, 'slot_index': 3}], 'outputs': [{'name': 'LATENT', 'type': 'LATENT', 'links': [6], 'shape': 3, 'slot_index': 0}], 'properties': {'Node name for S&R': 'KSampler'}, 'widgets_values': [583656379821656, 'randomize', 20, 8, 'euler', 'normal', 1]}, {'id': 11, 'type': 'EmptyLatentImage', 'pos': [-265, 183], 'size': {'0': 315, '1': 106}, 'flags': {}, 'order': 2, 'mode': 0, 'outputs': [{'name': 'LATENT', 'type': 'LATENT', 'links': [12], 'shape': 3}], 'properties': {'Node name for S&R': 'EmptyLatentImage'}, 'widgets_values': [512, 512, 1]}, {'id': 4, 'type': 'AnyNode', 'pos': [650, -185], 'size': {'0': 400, '1': 200}, 'flags': {}, 'order': 6, 'mode': 0, 'inputs': [{'name': 'any', 'type': '*', 'link': 6}, {'name': 'any2', 'type': '*', 'link': 5}], 'outputs': [{'name': 'any', 'type': '*', 'links': [1, 13], 'shape': 3, 'slot_index': 0}], 'properties': {'Node name for S&R': 'AnyNode'}, 'widgets_values': ["Output the type and some information about the inputs.\nOutput should be a string.\n\nOutput should have class information.\nIf it's a dict, then output the dict keys.\nIf it's a Tensor, output the shape.\netc.", 'gpt-4o']}, {'id': 13, 'type': 'AnyNodeCodeViewer', 'pos': [734, 129], 'size': {'0': 210, '1': 26}, 'flags': {}, 'order': 8, 'mode': 0, 'inputs': [{'name': 'ctrl', 'type': 'DICT', 'link': 13}], 'outputs': [{'name': 'DICT', 'type': 'DICT', 'links': None, 'shape': 3, 'slot_index': 0}], 'properties': {'Node name for S&R': 'AnyNodeCodeViewer'}}], 'links': [[1, 4, 0, 5, 0, 'STRING'], [5, 3, 0, 4, 1, '*'], [6, 7, 0, 4, 0, '*'], [7, 8, 0, 7, 0, 'MODEL'], [8, 8, 1, 9, 0, 'CLIP'], [9, 8, 1, 10, 0, 'CLIP'], [10, 9, 0, 7, 1, 'CONDITIONING'], [11, 10, 0, 7, 2, 'CONDITIONING'], [12, 11, 0, 7, 3, 'LATENT'], [13, 4, 0, 13, 0, 'DICT']], 'groups': [], 'config': {}, 'extra': {}, 'version': 0.4}
node_id = 13 # Example node ID
node_aware = NodeAware(workflow_data)
connections = node_aware.get_connected_nodes(node_id)
summary = node_aware.summarize_connections(node_id)
print("Connections:", connections)
print("Summary:\n", summary)
found_nodes = node_aware.find_node(id=node_id)
print("Found Node:", found_nodes)
+1 -1
View File
@@ -448,7 +448,7 @@
"Node name for S&R": "AnyNode"
},
"widgets_values": [
"I want you to output the image with an hue rotation by random degrees between 0 and 359 and tweaks the saturation and lightness drastically.\n\nUse the current system time as the random seed.\nUse torch where efficient.\n\nThe input should have four dimensions, (batch, width, height, color_channel) representing a batch of RGB images.\n\n\nHere is a working example that rotates hue by 180 and more importantly, outputs the right tensor shape:\n\n```python\ndef generated_function(input_data):\n # Convert input RGB tensor to a numpy array and normalize\n rgb_array = input_data.numpy()\n\n # Convert from RGB to HSV\n hsv_array = torch.empty_like(input_data)\n for i in range(rgb_array.shape[1]):\n for j in range(rgb_array.shape[2]):\n r, g, b = rgb_array[0,i,j] / 255.0\n max_c = max(r, g, b)\n min_c = min(r, g, b)\n delta = max_c - min_c\n \n # Calculate Hue\n if delta == 0:\n h = 0\n elif max_c == r:\n h = (60 * ((g - b) / delta) + 360) % 360\n elif max_c == g:\n h = (60 * ((b - r) / delta) + 120) % 360\n elif max_c == b:\n h = (60 * ((r - g) / delta) + 240) % 360\n \n # Calculate Saturation\n if max_c == 0:\n s = 0\n else:\n s = (delta / max_c)\n \n # Value is equal to max of R, G, B\n v = max_c\n \n # Shift Hue by 180 degrees\n h = (h + 180) % 360\n \n # Convert back to RGB\n c = v * s\n x = c * (1 - abs((h / 60) % 2 - 1))\n m = v - c\n \n if 0 <= h < 60:\n r1, g1, b1 = c, x, 0\n elif 60 <= h < 120:\n r1, g1, b1 = x, c, 0\n elif 120 <= h < 180:\n r1, g1, b1 = 0, c, x\n elif 180 <= h < 240:\n r1, g1, b1 = 0, x, c\n elif 240 <= h < 300:\n r1, g1, b1 = x, 0, c\n elif 300 <= h < 360:\n r1, g1, b1 = c, 0, x\n \n r, g, b = (r1 + m) * 255, (g1 + m) * 255, (b1 + m) * 255\n \n hsv_array[0,i,j,0] = r\n hsv_array[0,i,j,1] = g\n hsv_array[0,i,j,2] = b\n\n return hsv_array\n```"
"I want you to output the image with an hue rotation by random degrees between 0 and 359 and tweaks the saturation and lightness drastically.\n\nUse the current system time as the random seed. Output should be a tensor the same shape as input."
],
"color": "#2a363b",
"bgcolor": "#3f5159"