function registry first commit (barebones)

This commit is contained in:
holonic
2024-05-30 16:18:24 +01:00
parent bcac104e64
commit 704a4e649a
3 changed files with 68 additions and 34 deletions
+19 -11
View File
@@ -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
View File
@@ -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)
+1
View File
@@ -3,3 +3,4 @@ torch
numpy
scikit-learn
google-generativeai
chromadb