From 862bf0e132b3de0d4b9d62d5314023bc882b615b Mon Sep 17 00:00:00 2001 From: henri123lemoine Date: Mon, 27 Mar 2023 00:59:35 -0400 Subject: [PATCH] Commented out informed_assistant and semantic_search for testing. --- web/api/informed_assistant.py | 480 +++++++++++++++++----------------- web/api/requirements.txt | 1 + web/api/semantic_search.py | 248 +++++++++--------- web/package-lock.json | 2 +- web/package.json | 2 +- web/src/pages/index.tsx | 2 +- 6 files changed, 370 insertions(+), 365 deletions(-) diff --git a/web/api/informed_assistant.py b/web/api/informed_assistant.py index df1e49c..7435086 100644 --- a/web/api/informed_assistant.py +++ b/web/api/informed_assistant.py @@ -1,300 +1,300 @@ -# ---------------------------------- web code ---------------------------------- +# # ---------------------------------- web code ---------------------------------- -import json +# import json -from http.server import BaseHTTPRequestHandler +# from http.server import BaseHTTPRequestHandler -class handler(BaseHTTPRequestHandler): +# class handler(BaseHTTPRequestHandler): - # post request = calculate factorial of passed number - def do_POST(self): - self.send_response(200) - self.send_header('Content-type', 'application/json') - self.end_headers() - content_length = int(self.headers['Content-Length']) - post_data = self.rfile.read(content_length) - data = json.loads(post_data) +# # post request = calculate factorial of passed number +# def do_POST(self): +# self.send_response(200) +# self.send_header('Content-type', 'application/json') +# self.end_headers() +# content_length = int(self.headers['Content-Length']) +# post_data = self.rfile.read(content_length) +# data = json.loads(post_data) - results = {} +# results = {} - for i, link in enumerate(informed_assistant(data['query'])): - results[i] = json.dumps(link.__dict__) +# for i, link in enumerate(informed_assistant(data['query'])): +# results[i] = json.dumps(link.__dict__) - self.wfile.write(json.dumps(results).encode('utf-8')) +# self.wfile.write(json.dumps(results).encode('utf-8')) -# -------------------------------- non-web-code -------------------------------- -import time -import os -import openai +# # -------------------------------- non-web-code -------------------------------- +# import time +# import os +# import openai -import requests -from typing import List, Dict -import openai -import tiktoken -import asyncio +# import requests +# from typing import List, Dict +# import openai +# import tiktoken +# import asyncio -import config -from semantic_search import get_top_k_blocks +# import config +# from semantic_search import get_top_k_blocks -# OpenAI API key -try: - import config - OPENAI_API_KEY = config.OPENAI_API_KEY -except ImportError: - OPENAI_API_KEY = os.environ.get('OPENAI_API_KEY') -openai.api_key = OPENAI_API_KEY +# # OpenAI API key +# try: +# import config +# OPENAI_API_KEY = config.OPENAI_API_KEY +# except ImportError: +# OPENAI_API_KEY = os.environ.get('OPENAI_API_KEY') +# openai.api_key = OPENAI_API_KEY -# OpenAI models -EMBEDDING_MODEL = "text-embedding-ada-002" -COMPLETIONS_MODEL = "gpt-3.5-turbo" +# # OpenAI models +# EMBEDDING_MODEL = "text-embedding-ada-002" +# COMPLETIONS_MODEL = "gpt-3.5-turbo" -# OpenAI parameters -LEN_EMBEDDINGS = 1536 -MAX_LEN_PROMPT = 4095 # This may be 8191, unsure. +# # OpenAI parameters +# LEN_EMBEDDINGS = 1536 +# MAX_LEN_PROMPT = 4095 # This may be 8191, unsure. -# Paths -from pathlib import Path -project_path = Path(__file__).parent.parent.parent -PATH_TO_DATA = project_path / "web" / "api" / "data" / "alignment_texts.jsonl" # Path to the dataset .jsonl file. -PATH_TO_EMBEDDINGS = project_path / "web" / "api" / "data" / "embeddings.npy" # Path to the saved embeddings (.npy) file. -PATH_TO_DATASET = project_path / "web" / "api" / "data" / "dataset.pkl" # Path to the saved dataset (.pkl) file, containing the dataset class object. +# # Paths +# from pathlib import Path +# project_path = Path(__file__).parent.parent.parent +# PATH_TO_DATA = project_path / "web" / "api" / "data" / "alignment_texts.jsonl" # Path to the dataset .jsonl file. +# PATH_TO_EMBEDDINGS = project_path / "web" / "api" / "data" / "embeddings.npy" # Path to the saved embeddings (.npy) file. +# PATH_TO_DATASET = project_path / "web" / "api" / "data" / "dataset.pkl" # Path to the saved dataset (.pkl) file, containing the dataset class object. -class Dataset: - pass +# class Dataset: +# pass -class Block: - def __init__(self, title: str, author: str, date: str, url: str, tags: str, text: str): - self.title = title - self.author = author - self.date = date - self.url = url - self.tags = tags - self.text = text +# class Block: +# def __init__(self, title: str, author: str, date: str, url: str, tags: str, text: str): +# self.title = title +# self.author = author +# self.date = date +# self.url = url +# self.tags = tags +# self.text = text -MODERATION_ENDPOINT = "https://api.openai.com/v1/moderations" -def moderate_query(query: str) -> List[str]: - """This function uses the OpenAI Moderation API to check if a query contains any offensive language. +# MODERATION_ENDPOINT = "https://api.openai.com/v1/moderations" +# def moderate_query(query: str) -> List[str]: +# """This function uses the OpenAI Moderation API to check if a query contains any offensive language. - Args: - query (str): The query to be checked. +# Args: +# query (str): The query to be checked. - Raises: - Exception: If the API call fails. +# Raises: +# Exception: If the API call fails. - Returns: - List[str]: A list of categories that the query was flagged for. - """ +# Returns: +# List[str]: A list of categories that the query was flagged for. +# """ - headers = {"Content-Type": "application/json","Authorization": f"Bearer {OPENAI_API_KEY}"} +# headers = {"Content-Type": "application/json","Authorization": f"Bearer {OPENAI_API_KEY}"} - data = {"input": query} +# data = {"input": query} - response = requests.post(MODERATION_ENDPOINT, headers=headers, data=json.dumps(data)) - flagged_categories = [] +# response = requests.post(MODERATION_ENDPOINT, headers=headers, data=json.dumps(data)) +# flagged_categories = [] - if response.status_code == 200: - moderation_results = response.json() - flagged = moderation_results['results'][0]['flagged'] - categories = moderation_results['results'][0]['categories'] +# if response.status_code == 200: +# moderation_results = response.json() +# flagged = moderation_results['results'][0]['flagged'] +# categories = moderation_results['results'][0]['categories'] - if flagged: - for category, is_flagged in categories.items(): - if is_flagged: - flagged_categories.append(category) - else: - raise Exception(f"Error: {response.status_code} {response.reason}") +# if flagged: +# for category, is_flagged in categories.items(): +# if is_flagged: +# flagged_categories.append(category) +# else: +# raise Exception(f"Error: {response.status_code} {response.reason}") - return flagged_categories +# return flagged_categories -def limit_tokens(text: str, max_tokens: int, encoding_name: str = "cl100k_base") -> str: - encoding = tiktoken.get_encoding(encoding_name) - tokens = encoding.encode(text)[:max_tokens] - return encoding.decode(tokens) +# def limit_tokens(text: str, max_tokens: int, encoding_name: str = "cl100k_base") -> str: +# encoding = tiktoken.get_encoding(encoding_name) +# tokens = encoding.encode(text)[:max_tokens] +# return encoding.decode(tokens) -def generate_prompt(user_query: str, previous_dialogue: List[Dict[str, str]] = [], blocks: List[Block] = [], mode: str = "standard") -> List[Dict[str, str]]: - """ - This function generates a prompt in messages format for the OpenAI ChatCompletions API. - First, it picks a system description using the mode. - Second, it adds the previous dialogue to the prompt. - Third, it adds an instruction to the prompt based on the mode. - Fourth, it adds the context from the top-k most relevant blocks from the Alignment Research Dataset to the prompt. - Fifth, it adds the user query to the prompt. +# def generate_prompt(user_query: str, previous_dialogue: List[Dict[str, str]] = [], blocks: List[Block] = [], mode: str = "standard") -> List[Dict[str, str]]: +# """ +# This function generates a prompt in messages format for the OpenAI ChatCompletions API. +# First, it picks a system description using the mode. +# Second, it adds the previous dialogue to the prompt. +# Third, it adds an instruction to the prompt based on the mode. +# Fourth, it adds the context from the top-k most relevant blocks from the Alignment Research Dataset to the prompt. +# Fifth, it adds the user query to the prompt. - Messages take the following format: - messages=[ - {"role": "system", "content": "You are a helpful assistant."}, - {"role": "user", "content": "Who won the world series in 2020?"}, - {"role": "assistant", "content": "The Los Angeles Dodgers won the World Series in 2020."}, - {"role": "user", "content": "Where was it played?"} - ] +# Messages take the following format: +# messages=[ +# {"role": "system", "content": "You are a helpful assistant."}, +# {"role": "user", "content": "Who won the world series in 2020?"}, +# {"role": "assistant", "content": "The Los Angeles Dodgers won the World Series in 2020."}, +# {"role": "user", "content": "Where was it played?"} +# ] - Args: - user_query (str): The user query. - previous_dialogue (List[Dict[str, str]]): The previous dialogue. Defaults to []. - blocks (List[Block]): The top-k most relevant blocks from the Alignment Research Dataset. Defaults to []. - mode (str): The mode of the assistant. Can be "standard", etc. Defaults to "standard". +# Args: +# user_query (str): The user query. +# previous_dialogue (List[Dict[str, str]]): The previous dialogue. Defaults to []. +# blocks (List[Block]): The top-k most relevant blocks from the Alignment Research Dataset. Defaults to []. +# mode (str): The mode of the assistant. Can be "standard", etc. Defaults to "standard". - Returns: - List[Dict[str, str]]: The prompt in messages format. - """ - # Initialize prompt - prompt = [] +# Returns: +# List[Dict[str, str]]: The prompt in messages format. +# """ +# # Initialize prompt +# prompt = [] - # Generate system description - if mode == "standard": - prompt.append({"role": "system", "content": "You are a helpful assistant knowledgeable about AI Alignment and Safety."}) - # elif mode == "other": - else: - raise Exception(f"Invalid mode: {mode}") +# # Generate system description +# if mode == "standard": +# prompt.append({"role": "system", "content": "You are a helpful assistant knowledgeable about AI Alignment and Safety."}) +# # elif mode == "other": +# else: +# raise Exception(f"Invalid mode: {mode}") - # Add previous dialogue - for message in previous_dialogue: - prompt.append(message) +# # Add previous dialogue +# for message in previous_dialogue: +# prompt.append(message) - # Add instruction - if mode == "standard": - instruction_prompt = "Please answer my question (after the Q:) using the provided context." - prompt.append({"role": "assistant", "content": instruction_prompt}) - # elif mode == "other": - else: - raise Exception(f"Invalid mode: {mode}") +# # Add instruction +# if mode == "standard": +# instruction_prompt = "Please answer my question (after the Q:) using the provided context." +# prompt.append({"role": "assistant", "content": instruction_prompt}) +# # elif mode == "other": +# else: +# raise Exception(f"Invalid mode: {mode}") - # Add context from top-k blocks - if blocks is None: - return "Context missing." - context_prompt = "Context:\n\n" - for i, block in enumerate(blocks): - context_prompt += f"[{i}] {block.text}\n\n" - context_prompt = context_prompt[:-2] - context_prompt = limit_tokens(context_prompt, 2000) - prompt.append({"role": "user", "content": f"{context_prompt}"}) +# # Add context from top-k blocks +# if blocks is None: +# return "Context missing." +# context_prompt = "Context:\n\n" +# for i, block in enumerate(blocks): +# context_prompt += f"[{i}] {block.text}\n\n" +# context_prompt = context_prompt[:-2] +# context_prompt = limit_tokens(context_prompt, 2000) +# prompt.append({"role": "user", "content": f"{context_prompt}"}) - # Add user query - prompt.append({"role": "user", "content": f"Q: {user_query}"}) +# # Add user query +# prompt.append({"role": "user", "content": f"Q: {user_query}"}) - return prompt +# return prompt -def normal_completion(prompt: List[Dict[str, str]]) -> str: - """ - This function uses the OpenAI ChatCompletions API to answer a user query. +# def normal_completion(prompt: List[Dict[str, str]]) -> str: +# """ +# This function uses the OpenAI ChatCompletions API to answer a user query. - Args: - messages (Dict[str, str]): A dictionary containing the system prompt and user prompt, in addition to any previous dialogue. +# Args: +# messages (Dict[str, str]): A dictionary containing the system prompt and user prompt, in addition to any previous dialogue. - Returns: - str: The answer generated by the API. +# Returns: +# str: The answer generated by the API. - Raises: - Exception: If the API call fails. - """ - try: - return openai.ChatCompletion.create( - model=COMPLETIONS_MODEL, - messages=prompt - )["choices"][0]["message"]["content"] - except Exception as e: - print(e) - return "I'm sorry, I failed to process your query. Please try again. If the problem persists, please contact the administrator." +# Raises: +# Exception: If the API call fails. +# """ +# try: +# return openai.ChatCompletion.create( +# model=COMPLETIONS_MODEL, +# messages=prompt +# )["choices"][0]["message"]["content"] +# except Exception as e: +# print(e) +# return "I'm sorry, I failed to process your query. Please try again. If the problem persists, please contact the administrator." -async def stream_completion(prompt: List[Dict[str, str]], stream_delay: float = 0.1) -> str: - """ - This function uses the OpenAI ChatCompletions API to answer a user query, streaming the response. +# async def stream_completion(prompt: List[Dict[str, str]], stream_delay: float = 0.1) -> str: +# """ +# This function uses the OpenAI ChatCompletions API to answer a user query, streaming the response. - Args: - messages (Dict[str, str]): A dictionary containing the system prompt and user prompt, in addition to any previous dialogue. +# Args: +# messages (Dict[str, str]): A dictionary containing the system prompt and user prompt, in addition to any previous dialogue. - Returns: - str: The answer generated by the API. +# Returns: +# str: The answer generated by the API. - Raises: - Exception: If the API call fails. - """ - try: - async for part in await openai.ChatCompletion.acreate( - model=COMPLETIONS_MODEL, - messages=prompt, - stream=True - ): - finish_reason = part["choices"][0]["finish_reason"] - if "content" in part["choices"][0]["delta"]: - content = part["choices"][0]["delta"]["content"] - yield content - elif finish_reason: - print(f"Stream finished: {finish_reason}") - break - except Exception as e: - print(e) - response = "I'm sorry, I failed to process your query. Please try again. If the problem persists, please contact the administrator." - for word in response.split(): - time.sleep(stream_delay) - yield f"{word} " +# Raises: +# Exception: If the API call fails. +# """ +# try: +# async for part in await openai.ChatCompletion.acreate( +# model=COMPLETIONS_MODEL, +# messages=prompt, +# stream=True +# ): +# finish_reason = part["choices"][0]["finish_reason"] +# if "content" in part["choices"][0]["delta"]: +# content = part["choices"][0]["delta"]["content"] +# yield content +# elif finish_reason: +# print(f"Stream finished: {finish_reason}") +# break +# except Exception as e: +# print(e) +# response = "I'm sorry, I failed to process your query. Please try again. If the problem persists, please contact the administrator." +# for word in response.split(): +# time.sleep(stream_delay) +# yield f"{word} " -def informed_assistant(user_query: str, previous_dialogue: List[Dict[str, str]] = [], k: str = 10, mode: str = "standard", HyDE: bool = False, stream: bool = True, stream_delay: float = 0.1) -> str: - """ - This function uses the OpenAI ChatCompletions API to answer a user query. - It first checks if the query is offensive, and if so, raises an exception. - Then, it finds the top-k most relevant blocks from the Alignment Research Dataset and uses them as context for the ChatCompletions API. - It uses the blocks to generate a prompt for the ChatCompletions API. - Finally, it uses the ChatCompletions API to generate an answer to the user query. +# def informed_assistant(user_query: str, previous_dialogue: List[Dict[str, str]] = [], k: str = 10, mode: str = "standard", HyDE: bool = False, stream: bool = True, stream_delay: float = 0.1) -> str: +# """ +# This function uses the OpenAI ChatCompletions API to answer a user query. +# It first checks if the query is offensive, and if so, raises an exception. +# Then, it finds the top-k most relevant blocks from the Alignment Research Dataset and uses them as context for the ChatCompletions API. +# It uses the blocks to generate a prompt for the ChatCompletions API. +# Finally, it uses the ChatCompletions API to generate an answer to the user query. - Args: - user_query (str): The user query. - previous_dialogue (List[Dict[str, str]]): The previous dialogue. Defaults to []. - k (str): The number of blocks to use as context. - mode (str): The mode to use for the ChatCompletions API. Defaults to "standard". - HyDE (bool): Whether to use the HyDE technique for semantic search. This makes search slower, but better. Defaults to False. - stream (bool): Whether to stream the results from the ChatCompletions API. Defaults to True. - stream_delay (float): The delay between each word in the streamed response when streaming a hard-coded response. Defaults to 0.1. +# Args: +# user_query (str): The user query. +# previous_dialogue (List[Dict[str, str]]): The previous dialogue. Defaults to []. +# k (str): The number of blocks to use as context. +# mode (str): The mode to use for the ChatCompletions API. Defaults to "standard". +# HyDE (bool): Whether to use the HyDE technique for semantic search. This makes search slower, but better. Defaults to False. +# stream (bool): Whether to stream the results from the ChatCompletions API. Defaults to True. +# stream_delay (float): The delay between each word in the streamed response when streaming a hard-coded response. Defaults to 0.1. - Returns: - str: The answer to the user query. +# Returns: +# str: The answer to the user query. - Raises: - Exception: If the query is offensive. - """ - # 1. Check if the query is offensive - flagged_categories: List[str] = moderate_query(user_query) - if len(flagged_categories) > 0: - response = f"Your query contains offensive language. Please try again." - if stream: - for word in response.split(): - time.sleep(stream_delay) - yield f"{word} " - else: - return response +# Raises: +# Exception: If the query is offensive. +# """ +# # 1. Check if the query is offensive +# flagged_categories: List[str] = moderate_query(user_query) +# if len(flagged_categories) > 0: +# response = f"Your query contains offensive language. Please try again." +# if stream: +# for word in response.split(): +# time.sleep(stream_delay) +# yield f"{word} " +# else: +# return response - # 2. Find the top-k most relevant blocks from the Alignment Research Dataset - top_k_blocks: List[Block] = get_top_k_blocks(user_query, k, HyDE) +# # 2. Find the top-k most relevant blocks from the Alignment Research Dataset +# top_k_blocks: List[Block] = get_top_k_blocks(user_query, k, HyDE) - # 3. Generate a prompt for the ChatCompletions API - prompt: List[Dict[str, str]] = generate_prompt(user_query, previous_dialogue, top_k_blocks, mode) +# # 3. Generate a prompt for the ChatCompletions API +# prompt: List[Dict[str, str]] = generate_prompt(user_query, previous_dialogue, top_k_blocks, mode) - # 4. Use the top-k most relevant blocks as context for the ChatCompletions API, and generate an answer to the user query - if stream: - return stream_completion(prompt) - else: - return normal_completion(prompt) +# # 4. Use the top-k most relevant blocks as context for the ChatCompletions API, and generate an answer to the user query +# if stream: +# return stream_completion(prompt) +# else: +# return normal_completion(prompt) -if __name__ == "__main__": - # Test the question answering function - user_query = "Within the area of mitigating AI risk, there are several broad classes of action being taken. What does Technical safety research focus on?" - previous_dialogue = [ - {"role": "assistant", "content": "Hi! I know all about AI Alignment. Ask me a question!"}, - ] - k = 10 - mode = "standard" - HyDE = True - stream = False # Doesn't quite work yet +# if __name__ == "__main__": +# # Test the question answering function +# user_query = "Within the area of mitigating AI risk, there are several broad classes of action being taken. What does Technical safety research focus on?" +# previous_dialogue = [ +# {"role": "assistant", "content": "Hi! I know all about AI Alignment. Ask me a question!"}, +# ] +# k = 10 +# mode = "standard" +# HyDE = True +# stream = False # Doesn't quite work yet - print(asyncio.run(informed_assistant(user_query, previous_dialogue, k, mode, HyDE, stream))) +# print(asyncio.run(informed_assistant(user_query, previous_dialogue, k, mode, HyDE, stream))) - # if stream: - # for part in informed_assistant(user_query, previous_dialogue, k, mode, HyDE, stream): - # print(part, end="") - # else: - # print(informed_assistant(user_query, previous_dialogue, k, mode, HyDE, stream)) \ No newline at end of file +# # if stream: +# # for part in informed_assistant(user_query, previous_dialogue, k, mode, HyDE, stream): +# # print(part, end="") +# # else: +# # print(informed_assistant(user_query, previous_dialogue, k, mode, HyDE, stream)) \ No newline at end of file diff --git a/web/api/requirements.txt b/web/api/requirements.txt index 915b1bc..bf95cbd 100644 --- a/web/api/requirements.txt +++ b/web/api/requirements.txt @@ -1,3 +1,4 @@ openai==0.27.2 numpy==1.24.2 tenacity==8.2.2 +# aiohttp==3.8.3 \ No newline at end of file diff --git a/web/api/semantic_search.py b/web/api/semantic_search.py index e740f8c..d20dfc5 100644 --- a/web/api/semantic_search.py +++ b/web/api/semantic_search.py @@ -4,9 +4,12 @@ import json from http.server import BaseHTTPRequestHandler +def get_top_k_blocks(user_query: str, k: int = 10, HyDE: bool = False): + return "Hello World" + + class handler(BaseHTTPRequestHandler): - # post request = calculate factorial of passed number def do_POST(self): self.send_response(200) self.send_header('Content-type', 'application/json') @@ -24,153 +27,154 @@ class handler(BaseHTTPRequestHandler): # -------------------------------- non-web-code -------------------------------- -import time -import pickle -import os +# import time +# import pickle +# import os -import numpy as np +# import numpy as np -import openai -from openai.error import RateLimitError +# import openai +# from openai.error import RateLimitError -from functools import wraps -from typing import Callable, List, Type +# from functools import wraps +# from typing import Callable, List, Type -# OpenAI API key -try: - import config - openai.api_key = config.OPENAI_API_KEY -except ImportError: - openai.api_key = os.environ.get('OPENAI_API_KEY') +# # OpenAI API key +# try: +# import config +# openai.api_key = config.OPENAI_API_KEY +# except ImportError: +# openai.api_key = os.environ.get('OPENAI_API_KEY') -# OpenAI models -EMBEDDING_MODEL = "text-embedding-ada-002" -COMPLETIONS_MODEL = "gpt-3.5-turbo" +# # OpenAI models +# EMBEDDING_MODEL = "text-embedding-ada-002" +# COMPLETIONS_MODEL = "gpt-3.5-turbo" -# OpenAI parameters -LEN_EMBEDDINGS = 1536 -MAX_LEN_PROMPT = 4095 # This may be 8191, unsure. +# # OpenAI parameters +# LEN_EMBEDDINGS = 1536 +# MAX_LEN_PROMPT = 4095 # This may be 8191, unsure. -# Paths -from pathlib import Path -project_path = Path(__file__).parent.parent.parent -PATH_TO_DATA = project_path / "web" / "api" / "data" / "alignment_texts.jsonl" # Path to the dataset .jsonl file. -PATH_TO_EMBEDDINGS = project_path / "web" / "api" / "data" / "embeddings.npy" # Path to the saved embeddings (.npy) file. -PATH_TO_DATASET = project_path / "web" / "api" / "data" / "dataset.pkl" # Path to the saved dataset (.pkl) file, containing the dataset class object. +# # Paths +# from pathlib import Path +# project_path = Path(__file__).parent.parent.parent +# PATH_TO_DATA = project_path / "web" / "api" / "data" / "alignment_texts.jsonl" # Path to the dataset .jsonl file. +# PATH_TO_EMBEDDINGS = project_path / "web" / "api" / "data" / "embeddings.npy" # Path to the saved embeddings (.npy) file. +# PATH_TO_DATASET = project_path / "web" / "api" / "data" / "dataset.pkl" # Path to the saved dataset (.pkl) file, containing the dataset class object. -class Dataset: - pass +# class Dataset: +# pass -class Block: - def __init__(self, title: str, author: str, date: str, url: str, tags: str, text: str): - self.title = title - self.author = author - self.date = date - self.url = url - self.tags = tags - self.text = text +# class Block: +# def __init__(self, title: str, author: str, date: str, url: str, tags: str, text: str): +# self.title = title +# self.author = author +# self.date = date +# self.url = url +# self.tags = tags +# self.text = text -def retry_on_exception_types(exception_types: List[Type[Exception]], stop_after_attempt: int, max_wait_time: int) -> Callable: - def decorator(func: Callable) -> Callable: - @wraps(func) - def wrapper(*args, **kwargs): - attempts = 0 - while attempts < stop_after_attempt: - try: - return func(*args, **kwargs) - except tuple(exception_types) as e: - if attempts + 1 == stop_after_attempt: - raise e - wait_time = min(max_wait_time, (2 ** attempts)) # Exponential backoff - time.sleep(wait_time) - attempts += 1 - return wrapper - return decorator +# def retry_on_exception_types(exception_types: List[Type[Exception]], stop_after_attempt: int, max_wait_time: int) -> Callable: +# def decorator(func: Callable) -> Callable: +# @wraps(func) +# def wrapper(*args, **kwargs): +# attempts = 0 +# while attempts < stop_after_attempt: +# try: +# return func(*args, **kwargs) +# except tuple(exception_types) as e: +# if attempts + 1 == stop_after_attempt: +# raise e +# wait_time = min(max_wait_time, (2 ** attempts)) # Exponential backoff +# time.sleep(wait_time) +# attempts += 1 +# return wrapper +# return decorator -@retry_on_exception_types(exception_types=[RateLimitError], stop_after_attempt=4, max_wait_time=10) -def get_embedding(text: str) -> np.ndarray: - """Get the embedding for a given text. The wrapper function will retry with exponential backoffthe request if the API rate limit is reached, up to 4 times. +# @retry_on_exception_types(exception_types=[RateLimitError], stop_after_attempt=4, max_wait_time=10) +# def get_embedding(text: str) -> np.ndarray: +# """Get the embedding for a given text. The wrapper function will retry with exponential backoffthe request if the API rate limit is reached, up to 4 times. - Args: - text (str): The text to get the embedding for. +# Args: +# text (str): The text to get the embedding for. - Returns: - np.ndarray: The embedding for the given text. - """ - result = openai.Embedding.create( - model=EMBEDDING_MODEL, - input=text - ) - return result["data"][0]["embedding"] +# Returns: +# np.ndarray: The embedding for the given text. +# """ +# result = openai.Embedding.create( +# model=EMBEDDING_MODEL, +# input=text +# ) +# return result["data"][0]["embedding"] -def get_top_k_blocks(user_query: str, k: int = 10, HyDE: bool = False) -> List[Block]: - """Get the top k blocks that are most semantically similar to the query, using the provided dataset. +# def get_top_k_blocks(user_query: str, k: int = 10, HyDE: bool = False) -> List[Block]: +# """Get the top k blocks that are most semantically similar to the query, using the provided dataset. - Args: - query (str): The query to be searched for. - k (int, optional): The number of blocks to return. - HyDE (bool, optional): Whether to use HyDE or not. Defaults to False. +# Args: +# query (str): The query to be searched for. +# k (int, optional): The number of blocks to return. +# HyDE (bool, optional): Whether to use HyDE or not. Defaults to False. - Returns: - List[Block]: A list of the top k blocks that are most semantically similar to the query. - """ - # Get the dataset - with open(PATH_TO_DATASET, "rb") as f: - metadataset = pickle.load(f) +# Returns: +# List[Block]: A list of the top k blocks that are most semantically similar to the query. +# """ +# # Get the dataset +# with open(PATH_TO_DATASET, "rb") as f: +# metadataset = pickle.load(f) - # Get the embedding for the query. - query_embedding = get_embedding(user_query) +# # Get the embedding for the query. +# query_embedding = get_embedding(user_query) - # If HyDE is enabled, produce a no-context ChatCompletion to the query. - if HyDE: - messages = [ - {"role": "system", "content": "You are a knowledgeable AI Alignment assistant."}, - {"role": "user", "content": f"Do your best to answer the question/instruction, even if you don't know the correct answer or action for sure.\nQ: {user_query}"}, - ] - HyDE_completion = openai.ChatCompletion.create( - model=COMPLETIONS_MODEL, - messages=messages - )["choices"][0]["message"]["content"] - HyDe_completion_embedding = get_embedding(f"Question: {user_query}\n\nAnswer: {HyDE_completion}") +# # If HyDE is enabled, produce a no-context ChatCompletion to the query. +# if HyDE: +# messages = [ +# {"role": "system", "content": "You are a knowledgeable AI Alignment assistant."}, +# {"role": "user", "content": f"Do your best to answer the question/instruction, even if you don't know the correct answer or action for sure.\nQ: {user_query}"}, +# ] +# HyDE_completion = openai.ChatCompletion.create( +# model=COMPLETIONS_MODEL, +# messages=messages +# )["choices"][0]["message"]["content"] +# HyDe_completion_embedding = get_embedding(f"Question: {user_query}\n\nAnswer: {HyDE_completion}") - similarity_scores = np.dot(metadataset.embeddings, HyDe_completion_embedding) - else: - similarity_scores = np.dot(metadataset.embeddings, query_embedding) +# similarity_scores = np.dot(metadataset.embeddings, HyDe_completion_embedding) +# else: +# similarity_scores = np.dot(metadataset.embeddings, query_embedding) - ordered_blocks = np.argsort(similarity_scores)[::-1] # Sort the blocks by similarity score - top_k_text_indices = ordered_blocks[:k] # Get the top k indices of the blocks - top_k_metadata_indexes = [metadataset.embeddings_metadata_index[i] for i in top_k_text_indices] +# ordered_blocks = np.argsort(similarity_scores)[::-1] # Sort the blocks by similarity score +# top_k_text_indices = ordered_blocks[:k] # Get the top k indices of the blocks +# top_k_metadata_indexes = [metadataset.embeddings_metadata_index[i] for i in top_k_text_indices] - # Get the top k blocks (title, author, date, url, tags, text) - top_k_texts = [metadataset.embedding_strings[i] for i in top_k_text_indices] # Get the top k texts - top_k_metadata = [metadataset.metadata[i] for i in top_k_metadata_indexes] # Get the top k metadata (title, author, date, url, tags) +# # Get the top k blocks (title, author, date, url, tags, text) +# top_k_texts = [metadataset.embedding_strings[i] for i in top_k_text_indices] # Get the top k texts +# top_k_metadata = [metadataset.metadata[i] for i in top_k_metadata_indexes] # Get the top k metadata (title, author, date, url, tags) - print(f"Top {k} blocks for query: '{user_query}'") - print("=========================================") - print(f"Top_{k}_metadata: {top_k_metadata}") +# print(f"Top {k} blocks for query: '{user_query}'") +# print("=========================================") +# print(f"Top_{k}_metadata: {top_k_metadata}") - # Combine the top k texts and metadata into a list of Block objects - top_k_metadata_and_text = [list(top_k_metadata[i]) + [top_k_texts[i]] for i in range(k)] - top_k_blocks = [Block(*block) for block in top_k_metadata_and_text] +# # Combine the top k texts and metadata into a list of Block objects +# top_k_metadata_and_text = [list(top_k_metadata[i]) + [top_k_texts[i]] for i in range(k)] + +# top_k_blocks = [Block(*block) for block in top_k_metadata_and_text] - return top_k_blocks +# return top_k_blocks -if __name__ == "__main__": - # Test the embeddings function - query = "What is the best way to learn about AI alignment?" - k = 8 - HyDE = True +# if __name__ == "__main__": +# # Test the embeddings function +# query = "What is the best way to learn about AI alignment?" +# k = 8 +# HyDE = True - blocks = get_top_k_blocks(query, k, HyDE) - for block in blocks: - print(f"Title: {block.title}") - print(f"Author: {block.author}") - print(f"Date: {block.date}") - print(f"URL: {block.url}") - print(f"Tags: {block.tags}") - print(f"Text: {block.text}") - print() - print() \ No newline at end of file +# blocks = get_top_k_blocks(query, k, HyDE) +# for block in blocks: +# print(f"Title: {block.title}") +# print(f"Author: {block.author}") +# print(f"Date: {block.date}") +# print(f"URL: {block.url}") +# print(f"Tags: {block.tags}") +# print(f"Text: {block.text}") +# print() +# print() \ No newline at end of file diff --git a/web/package-lock.json b/web/package-lock.json index 4730a7c..1eb5ab9 100644 --- a/web/package-lock.json +++ b/web/package-lock.json @@ -8,7 +8,7 @@ "name": "alignment_search", "version": "0.1.0", "dependencies": { - "next": "^13.2.1", + "next": "^13.2.4", "react": "18.2.0", "react-dom": "18.2.0", "zod": "^3.20.6" diff --git a/web/package.json b/web/package.json index d9b988d..42e103f 100644 --- a/web/package.json +++ b/web/package.json @@ -9,7 +9,7 @@ "start": "next start" }, "dependencies": { - "next": "^13.2.1", + "next": "^13.2.4", "react": "18.2.0", "react-dom": "18.2.0", "zod": "^3.20.6" diff --git a/web/src/pages/index.tsx b/web/src/pages/index.tsx index 5011d39..eb29017 100644 --- a/web/src/pages/index.tsx +++ b/web/src/pages/index.tsx @@ -81,7 +81,7 @@ const SearchBox: React.FC = () => { onChange={(e) => setQuery(e.target.value)} />