NodeAware Update
This commit is contained in:
@@ -1,3 +1,4 @@
|
||||
__pycache__
|
||||
test
|
||||
notes
|
||||
.vscode
|
||||
|
||||
+383
-269
@@ -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"
|
||||
|
||||
@@ -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)
|
||||
@@ -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"
|
||||
|
||||
Reference in New Issue
Block a user