function registry first commit (barebones)
This commit is contained in:
+19
-11
@@ -13,7 +13,7 @@
|
||||
# - 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 os, json, random, string, sys, math, datetime, collections, itertools, functools, urllib, shutil, re, torch, time, decimal, matplotlib, io, base64, wave
|
||||
import os, json, random, string, sys, math, datetime, collections, itertools, functools, urllib, shutil, re, torch, time, decimal, matplotlib, io, base64, wave, chromadb
|
||||
import numpy
|
||||
import numpy as np
|
||||
import torch.nn.functional as F
|
||||
@@ -32,6 +32,7 @@ from openai import OpenAI
|
||||
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
|
||||
|
||||
# The template for the system message sent to ChatCompletions
|
||||
SYSTEM_TEMPLATE = """
|
||||
@@ -92,7 +93,7 @@ class AnyNode:
|
||||
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"}
|
||||
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"}
|
||||
|
||||
def __init__(self):
|
||||
self.script = None
|
||||
@@ -126,7 +127,8 @@ class AnyNode:
|
||||
|
||||
# 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.
|
||||
FUNCTION_CACHE = {}
|
||||
VERSION = "0.1.1"
|
||||
FUNCTION_REGISTRY = FunctionRegistry(schema="default", version=VERSION)
|
||||
|
||||
def render_template(self, template:str, any=None, seed=None):
|
||||
"""Render the system template with current state"""
|
||||
@@ -272,8 +274,12 @@ class AnyNode:
|
||||
"""Takes the prompt and inputs, Generates a function with an LLM for the Node"""
|
||||
result = None
|
||||
if not is_none(any):
|
||||
registry = self.FUNCTION_REGISTRY
|
||||
print(f"Last Error: {self.last_error}")
|
||||
if self.script is None or self.last_prompt != prompt or self.last_error is not None:
|
||||
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, **kwargs)
|
||||
@@ -281,7 +287,11 @@ class AnyNode:
|
||||
# Store the script for future use
|
||||
self.script = self.extract_imports(r)
|
||||
print(f"Stored script:\n{self.script}")
|
||||
self.last_prompt = prompt
|
||||
if use_function:
|
||||
self.script = fr['function']
|
||||
self.last_comment = fr['comment']
|
||||
self.imports = fr['imports']
|
||||
self.last_prompt = prompt
|
||||
|
||||
# Execute the stored script to define the function
|
||||
try:
|
||||
@@ -298,11 +308,7 @@ class AnyNode:
|
||||
print("--- Exception During Exec ---")
|
||||
# store the error for next run
|
||||
self.last_error = traceback.format_exc()
|
||||
if 'not defined' in self.last_error:
|
||||
# case where a library is missing
|
||||
raise e
|
||||
else:
|
||||
raise e
|
||||
raise e
|
||||
|
||||
# Assuming the generated code defines a function named 'generated_function'
|
||||
function_name = "generated_function"
|
||||
@@ -319,7 +325,9 @@ class AnyNode:
|
||||
else:
|
||||
print(f"Function '{function_name}' not found in generated code.")
|
||||
|
||||
self.last_error = None
|
||||
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)])
|
||||
return (result,)
|
||||
|
||||
class AnyNodeGemini(AnyNode):
|
||||
|
||||
+48
-23
@@ -6,15 +6,19 @@ from chromadb.config import Settings
|
||||
from chromadb.utils import embedding_functions
|
||||
|
||||
class FunctionRegistry:
|
||||
def __init__(self, registry_dir="output/anynode"):
|
||||
def __init__(self, registry_dir="output/anynode", schema="default", version="1.0"):
|
||||
self.registry_dir = registry_dir
|
||||
os.makedirs(self.registry_dir, exist_ok=True)
|
||||
self.registry_file = os.path.join(self.registry_dir, "function_registry.json")
|
||||
self.schema = schema
|
||||
self.version = version
|
||||
self.registry_file = os.path.join(self.registry_dir, f"function_registry_{self.schema}.json")
|
||||
self.registry = self.load_registry()
|
||||
self.chroma_client = self.init_chromadb()
|
||||
|
||||
def init_chromadb(self):
|
||||
settings = Settings(chroma_dir=os.path.join(self.registry_dir, "chroma_db"))
|
||||
folder = os.path.join(self.registry_dir, f"chroma_db_{self.schema}")
|
||||
print(f"ChromaDB Path: {folder}")
|
||||
settings = Settings(is_persistent=True, persist_directory=f"./{folder}")
|
||||
client = chromadb.Client(settings)
|
||||
return client
|
||||
|
||||
@@ -31,43 +35,60 @@ class FunctionRegistry:
|
||||
def hash_prompt(self, prompt):
|
||||
return hashlib.md5(prompt.encode('utf-8')).hexdigest()
|
||||
|
||||
def add_function(self, prompt, function_code):
|
||||
def add_function(self, prompt, function_code, imports, comment, input_types):
|
||||
prompt_hash = self.hash_prompt(prompt)
|
||||
self.registry[prompt_hash] = function_code
|
||||
self.registry[prompt_hash] = {
|
||||
"prompt": prompt,
|
||||
"imports": imports,
|
||||
"function": function_code,
|
||||
"comment": comment,
|
||||
"version": self.version
|
||||
}
|
||||
self.save_registry()
|
||||
self.add_function_to_chromadb(prompt_hash, prompt, function_code, imports, comment, input_types)
|
||||
|
||||
def get_function(self, prompt):
|
||||
prompt_hash = self.hash_prompt(prompt)
|
||||
return self.registry.get(prompt_hash, None)
|
||||
function_data = self.registry.get(prompt_hash, None)
|
||||
if function_data:
|
||||
return function_data
|
||||
return None
|
||||
|
||||
def query_chromadb(self, prompt, top_k=1):
|
||||
def query_chromadb(self, prompt, input_types, top_k=1):
|
||||
collection = self.chroma_client.get_or_create_collection(name="function_registry")
|
||||
results = collection.query(query_texts=[prompt], top_k=top_k)
|
||||
filters = {"input_types": input_types}
|
||||
results = collection.query(query_texts=[prompt], top_k=top_k, filter_metadata=filters)
|
||||
if results['documents']:
|
||||
return results['documents'][0]['content']
|
||||
return None
|
||||
|
||||
def add_function_to_chromadb(self, prompt, function_code):
|
||||
def add_function_to_chromadb(self, prompt_hash, prompt, function_code, imports, comment, input_types):
|
||||
collection = self.chroma_client.get_or_create_collection(name="function_registry")
|
||||
document = {
|
||||
"content": function_code,
|
||||
"metadata": {"prompt": prompt}
|
||||
metadata = {
|
||||
"prompt": prompt,
|
||||
"function": function_code,
|
||||
"imports": "\n".join(imports),
|
||||
"comment": comment,
|
||||
"input_types": "\n".join(input_types),
|
||||
"version": self.version
|
||||
}
|
||||
collection.add_documents([document])
|
||||
print(metadata)
|
||||
collection.add(
|
||||
documents=[prompt],
|
||||
metadatas=[metadata],
|
||||
ids=[prompt_hash]
|
||||
)
|
||||
#collection.add_documents([document])
|
||||
|
||||
def get_function_with_rag(self, prompt):
|
||||
def get_function_with_rag(self, prompt, input_types, top_k=1):
|
||||
function_code = self.get_function(prompt)
|
||||
if function_code is None:
|
||||
function_code = self.query_chromadb(prompt)
|
||||
function_code = self.query_chromadb(prompt, "\n".join(input_types), top_k=top_k)
|
||||
return function_code
|
||||
|
||||
def add_function_to_registry(self, prompt, function_code):
|
||||
self.add_function(prompt, function_code)
|
||||
self.add_function_to_chromadb(prompt, function_code)
|
||||
|
||||
if __name__ == "__main__":
|
||||
# Example Usage
|
||||
registry = FunctionRegistry()
|
||||
registry = FunctionRegistry(schema="default", version="1.0")
|
||||
|
||||
# Adding a function to the registry
|
||||
prompt = "Generate a function that multiplies the input by 5."
|
||||
@@ -75,8 +96,12 @@ if __name__ == "__main__":
|
||||
def generated_function(input_data):
|
||||
return input_data * 5
|
||||
"""
|
||||
registry.add_function_to_registry(prompt, function_code)
|
||||
imports = ["numpy", "math"]
|
||||
comment = "This function multiplies the input by 5."
|
||||
input_types = ["int"]
|
||||
registry.add_function(prompt, function_code, imports, comment, input_types)
|
||||
|
||||
# Retrieving a function from the registry
|
||||
retrieved_function = registry.get_function_with_rag(prompt)
|
||||
# Retrieving a function from the registry with top_k results
|
||||
top_k = 3
|
||||
retrieved_function = registry.get_function_with_rag(prompt, input_types, top_k=top_k)
|
||||
print("Retrieved Function:\n", retrieved_function)
|
||||
|
||||
@@ -3,3 +3,4 @@ torch
|
||||
numpy
|
||||
scikit-learn
|
||||
google-generativeai
|
||||
chromadb
|
||||
Reference in New Issue
Block a user