diff --git a/db.py b/db.py index 0f831ca..56bed76 100644 --- a/db.py +++ b/db.py @@ -1,10 +1,16 @@ import os from neo4j import GraphDatabase from dotenv import load_dotenv + load_dotenv() -URI = os.getenv("NEO4J_URI") -AUTH = (os.getenv('NEO4J_USERNAME'), os.getenv("NEO4J_PASSWORD")) +try: # to connect to the graph db + URI = os.environ["NEO4J_URI"] + assert URI + AUTH = (os.environ["NEO4J_USERNAME"], os.environ["NEO4J_PASSWORD"]) +except Exception as e: + raise ValueError("Error fetching Neo4J credentials from environment") from e + driver = GraphDatabase.driver(URI, auth=AUTH) driver.verify_connectivity() diff --git a/download_llama.py b/download_llama.py index 2bd6d68..9851056 100644 --- a/download_llama.py +++ b/download_llama.py @@ -6,7 +6,7 @@ import modal MODELS_DIR = "/llamas" # MODEL_NAME = "neuralmagic/Meta-Llama-3.1-70B-Instruct-quantized.w8a8" -MODEL_NAME = "neuralmagic/Meta-Llama-3.1-8B-Instruct-quantized.w8a16" +MODEL_NAME = "neuralmagic/Meta-Llama-3.1-8B-Instruct-quantized.w8a8" volume = modal.Volume.from_name("llamas", create_if_missing=True) @@ -59,4 +59,4 @@ def main( model_name: str = MODEL_NAME, force_download: bool = False, ): - download_model.remote(model_name, force_download) \ No newline at end of file + download_model.remote(model_name, force_download) diff --git a/expand.py b/expand.py index 0cf21bd..9e368ad 100644 --- a/expand.py +++ b/expand.py @@ -1,73 +1,91 @@ -import random -from typing import List, Literal, Optional +from typing import List, Literal from db import driver import structured_gen as sg from pydantic import BaseModel, Field from rich import print + # Define the data structures class Question(BaseModel): type: Literal["Question"] text: str + class Concept(BaseModel): type: Literal["Concept"] # Text must be lowercase - text: str = Field(pattern=r'^[a-z ]+$') + text: str = Field(pattern=r"^[a-z ]+$") + class ConceptWithLinks(Concept): relationship_type: Literal["IS_A", "AFFECTS", "CONNECTS_TO"] + class Answer(BaseModel): type: Literal["Answer"] text: str + # Permitted response formats class FromQuestion(BaseModel): """If at a question, may generate an answer.""" + answer: List[Answer] + class FromConcept(BaseModel): """If at a concept, may produce questions or relate to concepts""" + questions: List[Question] concepts: List[ConceptWithLinks] + class FromAnswer(BaseModel): """If at an answer, may generate concepts or new questions""" + concepts: List[Concept] questions: List[Question] + # Create a core node if it doesn't exist, or return the existing core node ID def get_or_make_core(question: str): with driver.session() as session: # Check if node exists and get its ID - result = session.run(""" - MATCH (n:Core {text: $question}) + result = session.run( + """ + MATCH (n:Core {text: $question}) RETURN n.id as id - """, question=question) + """, + question=question, + ) data = result.data() if len(data) > 0: return data[0]["id"] - + # Create new node with UUID if it doesn't exist - result = session.run(""" + result = session.run( + """ MERGE (n:Core {text: $question}) ON CREATE SET n.id = randomUUID() RETURN n.id as id - """, question=question) + """, + question=question, + ) data = result.data() if len(data) > 0: return data[0]["id"] else: raise ValueError(f"Failed to create new core node for question: {question}") - + + def load_neighbors(node_id: str, distance: int = 1): with driver.session() as session: - result = session.run(""" + result = session.run( + """ MATCH (node {id: $node_id})-[rel]-(neighbor) WHERE type(rel) <> "TRAVERSED" - RETURN + RETURN node.id as node_id, node.text as node_text, type(rel) as rel_type, @@ -75,17 +93,24 @@ def load_neighbors(node_id: str, distance: int = 1): neighbor.text as neighbor_text, labels(neighbor)[0] as neighbor_type, labels(node)[0] as node_type - """, node_id=node_id) + """, + node_id=node_id, + ) return result.data() - + + def load_node(node_id: str): with driver.session() as session: - result = session.run(""" + result = session.run( + """ MATCH (node {id: $node_id}) RETURN node.id as node_id, node.text as node_text, labels(node)[0] as label - """, node_id=node_id) + """, + node_id=node_id, + ) return result.single(strict=True) - + + # Node linking functions # Relationship Types (all have curiosity score 0-1): # RAISES -> (Concept/Core to Question) @@ -99,15 +124,22 @@ def question_to_concept(question: str, concept: str): except Exception as e: print(f"Skipping connection due to embedding error: {e}") return - + with driver.session() as session: - session.run(""" + session.run( + """ MERGE (question:Question {text: $question}) ON CREATE SET question.id = randomUUID(), question.embedding = $question_embedding MERGE (concept:Concept {text: $concept}) ON CREATE SET concept.id = randomUUID(), concept.embedding = $concept_embedding MERGE (concept)-[:RAISES]->(question) - """, question=question, concept=concept, question_embedding=question_embedding, concept_embedding=concept_embedding) + """, + question=question, + concept=concept, + question_embedding=question_embedding, + concept_embedding=concept_embedding, + ) + def question_to_answer(question: str, answer: str): try: @@ -116,24 +148,32 @@ def question_to_answer(question: str, answer: str): except Exception as e: print(f"Skipping connection due to embedding error: {e}") return - + with driver.session() as session: - session.run(""" + session.run( + """ MERGE (question:Question {text: $question}) ON CREATE SET question.id = randomUUID(), question.embedding = $question_embedding MERGE (answer:Answer {text: $answer}) ON CREATE SET answer.id = randomUUID(), answer.embedding = $answer_embedding MERGE (answer)-[:ANSWERS]->(question) - """, question=question, answer=answer, question_embedding=question_embedding, answer_embedding=answer_embedding) - + """, + question=question, + answer=answer, + question_embedding=question_embedding, + answer_embedding=answer_embedding, + ) + + def concept_to_concept( - concept1: str, - concept2: str, - relationship_type: Literal["IS_A", "AFFECTS", "CONNECTS_TO"]): + concept1: str, + concept2: str, + relationship_type: Literal["IS_A", "AFFECTS", "CONNECTS_TO"], +): # If the concepts are the same, don't create a relationship if concept1 == concept2: return - + try: concept1_embedding = sg.embed(concept1) concept2_embedding = sg.embed(concept2) @@ -149,7 +189,14 @@ def concept_to_concept( ON CREATE SET concept2.id = randomUUID(), concept2.embedding = $concept2_embedding MERGE (concept1)-[:{relationship_type}]->(concept2) """ - session.run(query, concept1=concept1, concept2=concept2, concept1_embedding=concept1_embedding, concept2_embedding=concept2_embedding) + session.run( + query, + concept1=concept1, + concept2=concept2, + concept1_embedding=concept1_embedding, + concept2_embedding=concept2_embedding, + ) + def concept_to_question(concept: str, question: str): try: @@ -160,13 +207,20 @@ def concept_to_question(concept: str, question: str): return with driver.session() as session: - session.run(""" + session.run( + """ MERGE (concept:Concept {text: $concept}) ON CREATE SET concept.id = randomUUID(), concept.embedding = $concept_embedding MERGE (question:Question {text: $question}) ON CREATE SET question.id = randomUUID(), question.embedding = $question_embedding MERGE (concept)-[:RAISES]->(question) - """, concept=concept, question=question, concept_embedding=concept_embedding, question_embedding=question_embedding) + """, + concept=concept, + question=question, + concept_embedding=concept_embedding, + question_embedding=question_embedding, + ) + # Core-specific functions def core_to_question(core: str, question: str): @@ -176,15 +230,22 @@ def core_to_question(core: str, question: str): except Exception as e: print(f"Skipping connection due to embedding error: {e}") return - + with driver.session() as session: - session.run(""" + session.run( + """ MERGE (core:Core {text: $core}) ON CREATE SET core.id = randomUUID(), core.embedding = $core_embedding MERGE (question:Question {text: $question}) ON CREATE SET question.id = randomUUID(), question.embedding = $question_embedding MERGE (core)-[:RAISES]->(question) - """, core=core, question=question, core_embedding=core_embedding, question_embedding=question_embedding) + """, + core=core, + question=question, + core_embedding=core_embedding, + question_embedding=question_embedding, + ) + def concept_to_core(concept: str, core: str): try: @@ -195,13 +256,20 @@ def concept_to_core(concept: str, core: str): return with driver.session() as session: - session.run(""" + session.run( + """ MERGE (concept:Concept {text: $concept}) ON CREATE SET concept.id = randomUUID(), concept.embedding = $concept_embedding MERGE (core:Core {text: $core}) ON CREATE SET core.id = randomUUID(), core.embedding = $core_embedding MERGE (concept)-[:EXPLAINS]->(core) - """, concept=concept, core=core, concept_embedding=concept_embedding, core_embedding=core_embedding) + """, + concept=concept, + core=core, + concept_embedding=concept_embedding, + core_embedding=core_embedding, + ) + def answer_to_concept(answer: str, concept: str): try: @@ -212,13 +280,20 @@ def answer_to_concept(answer: str, concept: str): return with driver.session() as session: - session.run(""" + session.run( + """ MERGE (answer:Answer {text: $answer}) ON CREATE SET answer.id = randomUUID(), answer.embedding = $answer_embedding MERGE (concept:Concept {text: $concept}) ON CREATE SET concept.id = randomUUID(), concept.embedding = $concept_embedding MERGE (answer)-[:SUGGESTS]->(concept) - """, answer=answer, concept=concept, answer_embedding=answer_embedding, concept_embedding=concept_embedding) + """, + answer=answer, + concept=concept, + answer_embedding=answer_embedding, + concept_embedding=concept_embedding, + ) + def answer_to_question(answer: str, question: str): try: @@ -227,43 +302,66 @@ def answer_to_question(answer: str, question: str): except Exception as e: print(f"Skipping connection due to embedding error: {e}") return - + with driver.session() as session: - session.run(""" + session.run( + """ MERGE (answer:Answer {text: $answer}) ON CREATE SET answer.id = randomUUID(), answer.embedding = $answer_embedding MERGE (question:Question {text: $question}) ON CREATE SET question.id = randomUUID(), question.embedding = $question_embedding MERGE (answer)-[:ANSWERS]->(question) - """, answer=answer, question=question, answer_embedding=answer_embedding, question_embedding=question_embedding) - -def record_traversal(from_node_id: str, to_node_id: str, traversal_type: Literal["random", "core", "neighbor"]): + """, + answer=answer, + question=question, + answer_embedding=answer_embedding, + question_embedding=question_embedding, + ) + + +def record_traversal( + from_node_id: str, + to_node_id: str, + traversal_type: Literal["random", "core", "neighbor"], +): with driver.session() as session: - session.run(""" + session.run( + """ MERGE (from_node {id: $from_node_id}) MERGE (to_node {id: $to_node_id}) MERGE (from_node)-[:TRAVERSED {timestamp: timestamp(), traversal_type: $traversal_type}]->(to_node) - """, from_node_id=from_node_id, to_node_id=to_node_id, traversal_type=traversal_type) - + """, + from_node_id=from_node_id, + to_node_id=to_node_id, + traversal_type=traversal_type, + ) + + def clear_db(): with driver.session() as session: - session.run(""" + session.run( + """ MATCH (n) DETACH DELETE n - """) + """ + ) + def random_node_id(): with driver.session() as session: - result = session.run(""" + result = session.run( + """ MATCH (n) RETURN n.id as id LIMIT 1 - """) + """ + ) return result.single(strict=True)["id"] + def format_node_neighborhood(node_id, truncate: bool = True): # Create ID mapping using ASCII uppercase letters (AA, AB, AC, etc.) id_counter = 0 uuid_to_simple_mapping = {} simple_to_uuid_mapping = {} - + def get_simple_id(): nonlocal id_counter # Generate IDs like AA, AB, ..., ZZ @@ -271,21 +369,21 @@ def format_node_neighborhood(node_id, truncate: bool = True): second = chr(65 + (id_counter % 26)) id_counter += 1 return f"NODE-{first}{second}" - + node = load_node(node_id) neighbors = load_neighbors(node_id) neighbors_string = f"{node['label'].upper()} {node['node_text']}\n" - + # Add direct neighbors if len(neighbors) > 0: neighbors_string += "\nDIRECT CONNECTIONS:\n" for neighbor in neighbors: - text = neighbor['neighbor_text'] + text = neighbor["neighbor_text"] if truncate: text = text[:70] + "..." if len(text) > 70 else text simple_id = get_simple_id() - simple_to_uuid_mapping[simple_id] = neighbor['neighbor_id'] - uuid_to_simple_mapping[neighbor['neighbor_id']] = simple_id + simple_to_uuid_mapping[simple_id] = neighbor["neighbor_id"] + uuid_to_simple_mapping[neighbor["neighbor_id"]] = simple_id neighbors_string += f"{simple_id:<8} {neighbor['rel_type']:<12} {neighbor['neighbor_type'].upper():<10} {text}\n" # Add semantically related nodes @@ -297,46 +395,54 @@ def format_node_neighborhood(node_id, truncate: bool = True): if nodes: # Only add section if there are related nodes neighbors_string += f"\n{node_type}s:\n" for n in nodes: - text = n['node_text'] + text = n["node_text"] if truncate: text = text[:70] + "..." if len(text) > 70 else text simple_id = get_simple_id() - simple_to_uuid_mapping[simple_id] = n['node_id'] - uuid_to_simple_mapping[n['node_id']] = simple_id + simple_to_uuid_mapping[simple_id] = n["node_id"] + uuid_to_simple_mapping[n["node_id"]] = simple_id neighbors_string += f"{simple_id:<8} {n['score']:<12.2f} {node_type.upper():<10} {text}\n" - + return neighbors_string, uuid_to_simple_mapping, simple_to_uuid_mapping + def find_related_nodes(node_id: str): with driver.session() as session: result = {} for node_type in ["Question", "Concept", "Answer"]: - result[node_type] = session.run(""" + result[node_type] = session.run( + """ MATCH (m {id: $node_id}) WHERE m.embedding IS NOT NULL CALL db.index.vector.queryNodes( - $vector_index_name, - $limit, + $vector_index_name, + $limit, m.embedding ) YIELD node, score RETURN node.id as node_id, node.text as node_text, score - """, node_id=node_id, vector_index_name=f"{node_type.lower()}_embedding", limit=10).data() + """, + node_id=node_id, + vector_index_name=f"{node_type.lower()}_embedding", + limit=10, + ).data() return result + def remove_index(index_name: str): with driver.session() as session: - session.run(f""" + session.run( + f""" DROP INDEX {index_name} IF EXISTS - """) + """ + ) -if __name__ == "__main__": - # Clear the database - # print("WARNING: Clearing the database") - # clear_db() - # Set the purpose - purpose = "Support humanity" +def main(do_clear_db=False, purpose="Support humanity"): + # Clear the database if requested + if do_clear_db: + print("WARNING: Clearing the database") + clear_db() # Create the core node and get its ID current_node_id = get_or_make_core(purpose) @@ -345,7 +451,6 @@ if __name__ == "__main__": # Get embedding dimensions embedding_dimensions = len(sg.embed(purpose)) - # Remove existing indices remove_index("core_id") remove_index("question_embedding") @@ -357,15 +462,16 @@ if __name__ == "__main__": # Create regular indices index_queries = [ "CREATE INDEX core_id IF NOT EXISTS FOR (n:Core) ON (n.id)", - "CREATE INDEX question_id IF NOT EXISTS FOR (n:Question) ON (n.id)", + "CREATE INDEX question_id IF NOT EXISTS FOR (n:Question) ON (n.id)", "CREATE INDEX concept_id IF NOT EXISTS FOR (n:Concept) ON (n.id)", - "CREATE INDEX answer_id IF NOT EXISTS FOR (n:Answer) ON (n.id)" + "CREATE INDEX answer_id IF NOT EXISTS FOR (n:Answer) ON (n.id)", ] - + # Create vector indices vector_index_queries = [] for node_type in ["Question", "Concept", "Answer"]: - vector_index_queries.append(f""" + vector_index_queries.append( + f""" CREATE VECTOR INDEX {node_type.lower()}_embedding IF NOT EXISTS FOR (n:{node_type}) ON (n.embedding) OPTIONS {{ @@ -374,7 +480,8 @@ if __name__ == "__main__": `vector.similarity_function`: 'COSINE' }} }} - """) + """ + ) # Execute all queries for query in index_queries + vector_index_queries: @@ -390,7 +497,10 @@ if __name__ == "__main__": # Get the user prompt. Shows previous nodes and actions, then # shows the current node. - prompt = "\n".join([f"{n['label'].upper()} {n['node_text']}" for n in history]) + f"\nCurrent node: {current_node_label.upper()} {current_node_text}" + prompt = ( + "\n".join([f"{n['label'].upper()} {n['node_text']}" for n in history]) + + f"\nCurrent node: {current_node_label.upper()} {current_node_text}" + ) prompt = "Here is the traversal history:\n" + prompt # prompt += f"Here are nodes related to the current node:\n" +\ # format_node_neighborhood(current_node_id, truncate=False) @@ -411,16 +521,16 @@ if __name__ == "__main__": # Get the system prompt system_prompt = f""" - You are a superintelligent AI building a self-expanding knowledge graph. + You are a superintelligent AI building a self-expanding knowledge graph. Your goal is to achieve the core directive "{purpose}". - + Generate an expansion of the current node. An expansion may include: - - A list of new questions. - - Questions should be short and direct. + - A list of new questions. + - Questions should be short and direct. - If you generate multiple questions, they should be distinct and not similar. - - A list of new concepts. - - Concepts are words or short combinations of words that + - A list of new concepts. + - Concepts are words or short combinations of words that are related to the current node. - Concepts may connect to each other. - Concepts may be related by IS_A, AFFECTS, or CONNECTS_TO. @@ -450,9 +560,11 @@ if __name__ == "__main__": result = sg.generate_by_schema( sg.messages(user=prompt, system=system_prompt), - result_format.model_json_schema() + result_format.model_json_schema(), + ) + expansion = result_format.model_validate_json( + result.choices[0].message.content ) - expansion = result_format.model_validate_json(result.choices[0].message.content) except Exception as e: print(f"Error generating expansion: {e}") # Return to the core node @@ -470,7 +582,9 @@ if __name__ == "__main__": for purpose in expansion.questions: concept_to_question(current_node_text, purpose.text) for concept in expansion.concepts: - concept_to_concept(current_node_text, concept.text, concept.relationship_type) + concept_to_concept( + current_node_text, concept.text, concept.relationship_type + ) # If we are at an answer, we can link to the concepts and questions. elif current_node_label == "Answer": @@ -490,25 +604,31 @@ if __name__ == "__main__": neighbors = load_neighbors(current_node_id) # Formatting the neighbor table - neighbors_string, uuid_to_simple_mapping, simple_to_uuid_mapping = format_node_neighborhood(current_node_id) + ( + neighbors_string, + uuid_to_simple_mapping, + simple_to_uuid_mapping, + ) = format_node_neighborhood(current_node_id) # Choose a new node if there are any neighbors if len(neighbors) > 0: old_node_id = current_node_id - print("----------------------------------------------------------------------------------") + print( + "----------------------------------------------------------------------------------" + ) print(neighbors_string) # Construct selectable nodes selectable_nodes = set() for neighbor in neighbors: # Add the neighbor's simple ID - selectable_nodes.add(uuid_to_simple_mapping[neighbor['neighbor_id']]) + selectable_nodes.add(uuid_to_simple_mapping[neighbor["neighbor_id"]]) # Add all the keys in the uuid_to_simple_mapping selectable_nodes.update(simple_to_uuid_mapping.keys()) - selectable_nodes.add('random') + selectable_nodes.add("random") # selectable_nodes.add('core') # Remove the current node from the selectable nodes if it's in there @@ -516,21 +636,23 @@ if __name__ == "__main__": if current_node_id in selectable_nodes: selectable_nodes.remove(current_node_id) - choice_prompt = prompt + \ - "Select a node to traverse to. Respond with the node ID." + \ - "You will generate a new expansion of the node you traverse to." + \ - "You will not be able to choose the current node." + \ - "You may choose 'random' to choose a random node." - # "You may also choose 'core' to return to the core node, " + \ - # "or 'random' to choose a random node." + choice_prompt = ( + prompt + + "Select a node to traverse to. Respond with the node ID." + + "You will generate a new expansion of the node you traverse to." + + "You will not be able to choose the current node." + + "You may choose 'random' to choose a random node." + ) + # "You may also choose 'core' to return to the core node, " + \ + # "or 'random' to choose a random node." node_selection = sg.choose( sg.messages(user=choice_prompt, system=system_prompt), - choices=list(selectable_nodes) + choices=list(selectable_nodes), ) - is_random = node_selection == 'random' - is_core = node_selection == 'core' + is_random = node_selection == "random" + is_core = node_selection == "core" if is_random: current_node_id = random_node_id() @@ -545,6 +667,33 @@ if __name__ == "__main__": print(f"SELECTED {node['label'].upper()} {node['node_text']}\n") history.append(current_node) - - traversal_type = 'random' if is_random else 'core' if is_core else 'neighbor' + + traversal_type = ( + "random" if is_random else "core" if is_core else "neighbor" + ) record_traversal(old_node_id, current_node_id, traversal_type) + + +if __name__ == "__main__": + import argparse + + parser = argparse.ArgumentParser( + description="Run a self-expanding knowledge graph around a core purpose." + ) + + parser.add_argument( + "--do-clear-db", + "--do_clear_db", + action="store_true", + help="If set, clear the database before proceeding.", + ) + parser.add_argument( + "--purpose", + type=str, + default="Support humanity", + help='Set the purpose (default: "Support humanity").', + ) + + args = parser.parse_args() + + main(args.do_clear_db, args.purpose) diff --git a/modal_embeddings.py b/modal_embeddings.py index 889983f..9016932 100644 --- a/modal_embeddings.py +++ b/modal_embeddings.py @@ -1,18 +1,18 @@ -import json -import os import socket import subprocess from pathlib import Path import modal -GPU_CONFIG = modal.gpu.A10G() +GPU_CONFIG = modal.gpu.L40S() MODEL_ID = "BAAI/bge-base-en-v1.5" BATCH_SIZE = 32 DOCKER_IMAGE = ( - "ghcr.io/huggingface/text-embeddings-inference:86-0.4.0" # Ampere 86 for A10s. - # "ghcr.io/huggingface/text-embeddings-inference:0.4.0" # Ampere 80 for A100s. - # "ghcr.io/huggingface/text-embeddings-inference:0.3.0" # Turing for T4s. + # "ghcr.io/huggingface/text-embeddings-inference:hopper-1.6" # Hopper 90 for H100s (marked experimental) + "ghcr.io/huggingface/text-embeddings-inference:89-1.6" # Lovelace 89 for L40S + # "ghcr.io/huggingface/text-embeddings-inference:86-0.4.0" # Ampere 86 for A10 + # "ghcr.io/huggingface/text-embeddings-inference:0.4.0" # Ampere 80 for A100 + # "ghcr.io/huggingface/text-embeddings-inference:0.3.0" # Turing for T4 ) DATA_PATH = Path("/data/dataset.jsonl") @@ -39,16 +39,15 @@ def spawn_server() -> subprocess.Popen: # If so, a connection can never be made. retcode = process.poll() if retcode is not None: - raise RuntimeError( - f"launcher exited unexpectedly with code {retcode}" - ) + raise RuntimeError(f"launcher exited unexpectedly with code {retcode}") def download_model(): # Wait for server to start. This downloads the model weights when not present. spawn_server().terminate() -app = modal.App("cameron-embeddings") + +app = modal.App("self-expansion-embeddings") tei_image = ( modal.Image.from_registry( @@ -63,6 +62,7 @@ tei_image = ( with tei_image.imports(): from httpx import AsyncClient + @app.cls( gpu=GPU_CONFIG, image=tei_image, @@ -91,9 +91,10 @@ class TextEmbeddingsInference: image = modal.Image.debian_slim(python_version="3.10").pip_install( - "pandas", "db-dtypes", "tqdm" + "pandas", "db-dtypes", "tqdm" ) + @app.function( image=image, ) @@ -102,7 +103,8 @@ def embed(data): return model.embed.remote(data)[0] + @app.local_entrypoint() -def main(): - embeddings = embed.remote([('hello', 'world'), ('hello', 'world')]) - print(embeddings) +def main(text: str = "hello"): + embedding = embed.remote([text]) + print(text, embedding[:10], "...") diff --git a/modal_vllm_container.py b/modal_vllm_container.py index 23e161c..950fb6b 100644 --- a/modal_vllm_container.py +++ b/modal_vllm_container.py @@ -1,13 +1,12 @@ import subprocess -from pathlib import Path import modal -from modal import App, Image, Mount, Secret, gpu +from modal import App, Image, Mount, gpu from download_llama import MODEL_NAME, MODELS_DIR ########## CONSTANTS ########## -MODEL_PATH = MODELS_DIR + '/' + MODEL_NAME +MODEL_PATH = MODELS_DIR + "/" + MODEL_NAME # define model for serving and path to store in modal container SECONDS = 60 # for timeout @@ -53,7 +52,7 @@ vllm_image = ( Image.debian_slim(python_version="3.12") .pip_install( [ - "vllm", + "vllm==0.6.6.post1", "huggingface_hub", "hf-transfer", "ray", @@ -67,7 +66,7 @@ vllm_image = ( ########## APP SETUP ########## -app = App("cameron-vllm") +app = App("self-expansion-vllm") NO_GPU = 1 TOKEN = "super-secret-token" # for demo purposes, for production, you can use Modal secrets to store token @@ -75,9 +74,10 @@ TOKEN = "super-secret-token" # for demo purposes, for production, you can use M # https://github.com/chujiezheng/chat_templates/tree/main/chat_templates LOCAL_TEMPLATE_PATH = "template_llama3.jinja" + @app.function( image=vllm_image, - gpu=gpu.A100(count=NO_GPU, size="80GB"), + gpu=gpu.L40S(count=NO_GPU), container_idle_timeout=20 * SECONDS, volumes={MODELS_DIR: volume}, mounts=[ @@ -85,9 +85,7 @@ LOCAL_TEMPLATE_PATH = "template_llama3.jinja" LOCAL_TEMPLATE_PATH, remote_path="/root/template_llama3.jinja" ) ], - # https://modal.com/docs/guide/concurrent-inputs - concurrency_limit=1, # fix at 1 to test concurrency within 1 server setup - allow_concurrent_inputs=256, # max concurrent input into container + allow_concurrent_inputs=256, # max concurrent input into container -- effectively batch size ) @modal.web_server(port=8000, startup_timeout=60 * SECONDS) def serve(): @@ -99,4 +97,31 @@ def serve(): --chat-template /root/template_llama3.jinja """ print(cmd) - subprocess.Popen(cmd, shell=True) \ No newline at end of file + subprocess.Popen(cmd, shell=True) + + +@app.function( + image=vllm_image, + gpu=gpu.L40S(count=NO_GPU), + container_idle_timeout=20 * SECONDS, + volumes={MODELS_DIR: volume}, + mounts=[ + Mount.from_local_file( + LOCAL_TEMPLATE_PATH, remote_path="/root/template_llama3.jinja" + ) + ], + # https://modal.com/docs/guide/concurrent-inputs +) +def infer(prompt: str = "How many r's are in the word 'strawberry'?"): + from vllm import LLM + + conversation = [{"role": "user", "content": prompt}] + + llm = LLM(model=MODEL_PATH, generation_config="auto") + response = llm.chat(conversation, use_tqdm=True)[0].outputs[0].text + return response + + +@app.local_entrypoint() +def main(prompt: str = "How many r's are in the word 'strawberry'?"): + print(infer.remote(prompt)) diff --git a/structured_gen.py b/structured_gen.py index 0c28d4b..59fd1db 100644 --- a/structured_gen.py +++ b/structured_gen.py @@ -5,6 +5,7 @@ from typing import List, Dict import os import dotenv + dotenv.load_dotenv() CLIENT = OpenAI( @@ -19,12 +20,14 @@ print("Using model:", DEFAULT_MODEL) MAX_TOKENS = 12000 + def messages(user: str, system: str = "You are a helpful assistant."): ms = [{"role": "user", "content": user}] if system: ms.insert(0, {"role": "system", "content": system}) return ms + def generate( messages: List[Dict[str, str]], response_format: BaseModel, @@ -32,15 +35,15 @@ def generate( response = CLIENT.beta.chat.completions.parse( model=DEFAULT_MODEL, messages=messages, - response_format=response_format, extra_body={ # 'guided_decoding_backend': 'outlines', "max_tokens": MAX_TOKENS, - } + }, ) return response + def generate_by_schema( messages: List[Dict[str, str]], schema: str, @@ -52,10 +55,11 @@ def generate_by_schema( # 'guided_decoding_backend': 'outlines', "max_tokens": MAX_TOKENS, "guided_json": schema, - } + }, ) return response + def choose( messages: List[Dict[str, str]], choices: List[str], @@ -67,6 +71,7 @@ def choose( ) return completion.choices[0].message.content + def regex( messages: List[Dict[str, str]], regex: str, @@ -78,6 +83,7 @@ def regex( ) return completion.choices[0].message.content + def embed(content: str) -> List[float]: - f = modal.Function.lookup("cameron-embeddings", "embed") + f = modal.Function.lookup("self-expansion-embeddings", "embed") return f.remote(content) -- 2.51.2 From bd7a59336a2aa218170527ab576eb03dc6277fba Mon Sep 17 00:00:00 2001 From: Charles Frye Date: Fri, 24 Jan 2025 16:46:48 -0800 Subject: [PATCH 2/3] add a sample environment file --- .env.example | 5 +++++ 1 file changed, 5 insertions(+) create mode 100644 .env.example diff --git a/.env.example b/.env.example new file mode 100644 index 0000000..8ceaf09 --- /dev/null +++ b/.env.example @@ -0,0 +1,5 @@ +NEO4J_URI= +NEO4J_USERNAME= +NEO4J_PASSWORD= +VLLM_BASE_URL=https://url--that-ends-in-serve.modal.run/v1 +VLLM_TOKEN=super-secret-token -- 2.51.2 From c8343f314ed905d0157091c2e1773e8f9e5348b3 Mon Sep 17 00:00:00 2001 From: Charles Frye Date: Fri, 24 Jan 2025 16:47:02 -0800 Subject: [PATCH 3/3] add a minimal README --- README.md | 15 +++++++++++++++ 1 file changed, 15 insertions(+) create mode 100644 README.md diff --git a/README.md b/README.md new file mode 100644 index 0000000..f9820dd --- /dev/null +++ b/README.md @@ -0,0 +1,15 @@ + # Self-Expanding Knowledge Graph + + +```bash +# set environment variables in .env as in .env.example +pip install -r requirements.txt +modal setup +modal run modal_embeddings.py # test run of embedding service +modal deploy modal_embeddings.py # deploy embedding service +modal run modal_vllm_container.py # test run of llm +modal deploy modal_vllm_container.py # deploy llm service +# now, look for a URL in the terminal that includes -serve.modal.run +# set that PLUS "/v1" at the end as VLLM_BASE_URL in .env +python expand.py --purpose "Do dogs know that their dreams aren't real?" # start expanding +```