From 5900da48bb381030e54e5abfdd2b7fcfff7e4c7f Mon Sep 17 00:00:00 2001 From: henri123lemoine Date: Tue, 27 Jun 2023 02:45:50 -0400 Subject: [PATCH] Renamed classes, added settings --- src/dataset/pinecone_db_handler.py | 2 +- src/dataset/settings.py | 19 ++++++++++++++++-- src/dataset/sql_db_handler.py | 16 +++++---------- src/dataset/update_dataset.py | 32 ++++++++++++------------------ src/main.py | 11 +++++----- 5 files changed, 41 insertions(+), 39 deletions(-) diff --git a/src/dataset/pinecone_db_handler.py b/src/dataset/pinecone_db_handler.py index 6f53daa..a0db67d 100644 --- a/src/dataset/pinecone_db_handler.py +++ b/src/dataset/pinecone_db_handler.py @@ -10,7 +10,7 @@ import logging logger = logging.getLogger(__name__) -class PineconeDB: +class PineconeDBHandler: def __init__( self, create_index: bool = False, diff --git a/src/dataset/settings.py b/src/dataset/settings.py index 892408a..f8045a3 100644 --- a/src/dataset/settings.py +++ b/src/dataset/settings.py @@ -1,6 +1,21 @@ +from pathlib import Path + +### FILE PATHS ### +current_file_path = Path(__file__).resolve() +SQL_DB_PATH = str(current_file_path.parent / 'data' / 'ARD.db') + +### DATASET ### +ARD_DATASET_NAME = "StampyAI/alignment-research-dataset" + +### EMBEDDINGS ### +EMBEDDINGS_MODEL = "text-embedding-ada-002" +EMBEDDING_DIMS = 1536 ### PINECONE ### PINECONE_INDEX_NAME = "stampy-chat-embeddings-test" -PINECONE_VALUES_DIMS = 1536 +PINECONE_VALUES_DIMS = EMBEDDING_DIMS PINECONE_METRIC = "cosine" -PINECONE_METADATA_ENTRIES = ["entry_id", "source", "title", "authors", "text"] \ No newline at end of file +PINECONE_METADATA_ENTRIES = ["entry_id", "source", "title", "authors", "text"] + +### MISC ### +CUSTOM_SOURCES = ['gwern_blog'] \ No newline at end of file diff --git a/src/dataset/sql_db_handler.py b/src/dataset/sql_db_handler.py index c2a891f..05afbf3 100644 --- a/src/dataset/sql_db_handler.py +++ b/src/dataset/sql_db_handler.py @@ -1,21 +1,15 @@ -import os import sqlite3 from typing import List, Dict, Any +from .settings import SQL_DB_PATH + import logging logger = logging.getLogger(__name__) -class DatabaseHandler: - def __init__( - self, - db_name: str = "data\\alignment_database.db", - ): - # Get the directory of this script - script_dir = os.path.dirname(os.path.realpath(__file__)) - - # Combine the script directory with the relative database path - self.db_name = os.path.join(script_dir, db_name) +class SQLDBHandler: + def __init__(self): + self.db_name = SQL_DB_PATH self.create_tables() diff --git a/src/dataset/update_dataset.py b/src/dataset/update_dataset.py index a8b5c49..d9d2e8d 100644 --- a/src/dataset/update_dataset.py +++ b/src/dataset/update_dataset.py @@ -7,8 +7,10 @@ import openai from datasets import load_dataset from .text_splitter import TokenSplitter -from .pinecone_handler import PineconeHandler -from .database_handler import DatabaseHandler +from .sql_db_handler import SQLDBHandler +from .pinecone_db_handler import PineconeDBHandler + +from .settings import EMBEDDINGS_MODEL, EMBEDDING_DIMS, ARD_DATASET_NAME import logging logger = logging.getLogger(__name__) @@ -20,21 +22,13 @@ class ARDUpdater: min_tokens_per_block: int = 200, # Minimum number of tokens per block. max_tokens_per_block: int = 400, # Maximum number of tokens per block. rate_limit_per_minute: int = 3_500, # Rate limit for the OpenAI API. - embedding_model="text-embedding-ada-002", - embedding_dims=1536, ): self.rate_limit_per_minute = rate_limit_per_minute self.delay_in_seconds = 60.0 / self.rate_limit_per_minute - self.embedding_model = embedding_model - self.embedding_dims = embedding_dims - - self.token_splitter = TokenSplitter( - min_tokens=min_tokens_per_block, - max_tokens=max_tokens_per_block - ) - self.db = DatabaseHandler() - self.pinecone_db = PineconeHandler() + self.token_splitter = TokenSplitter(min_tokens_per_block, max_tokens_per_block) + self.sql_db = SQLDBHandler() + self.pinecone_db = PineconeDBHandler() def update(self, custom_sources: List[str] = ['all']): for source in custom_sources: @@ -43,10 +37,10 @@ class ARDUpdater: def update_source(self, source: str): logger.info(f"Updating {source} entries...") - iterable_data = load_dataset('StampyAI/alignment-research-dataset', source, split='train', streaming=True) + iterable_data = load_dataset(ARD_DATASET_NAME, source, split='train', streaming=True) iterable_data = iterable_data.map(self.preprocess) iterable_data = iterable_data.filter(lambda entry: entry is not None) - iterable_data = iterable_data.filter(lambda entry: self.db.upsert_entry(entry)) + iterable_data = iterable_data.filter(lambda entry: self.sql_db.upsert_entry(entry)) for entry in tqdm(iterable_data): try: @@ -56,7 +50,7 @@ class ARDUpdater: chunks = self.token_splitter.split(entry['text'], signature) embeddings = self.get_embeddings(chunks) - self.db.upsert_chunks(entry['id'], chunks) + self.sql_db.upsert_chunks(entry['id'], chunks) self.pinecone_db.insert_entry(entry, chunks, embeddings) except Exception as e: logger.error(f"An error occurred while updating source {source}: {str(e)}", exc_info=True) @@ -99,10 +93,10 @@ class ARDUpdater: raise ValueError(f"Entry text is too short (< {len_lower_limit} tokens).") def get_embeddings(self, chunks): - embeddings = np.zeros((len(chunks), self.embedding_dims)) + embeddings = np.zeros((len(chunks), EMBEDDING_DIMS)) openai_output = openai.Embedding.create( - model=self.embedding_model, + model=EMBEDDINGS_MODEL, input=chunks )['data'] @@ -112,7 +106,7 @@ class ARDUpdater: return embeddings def reset_dbs(self): - self.db.create_tables(True) + self.sql_db.create_tables(True) self.pinecone_db.create_index(True) diff --git a/src/main.py b/src/main.py index 23a91e4..eb6f418 100644 --- a/src/main.py +++ b/src/main.py @@ -1,18 +1,16 @@ # main.py import os -import openai - -from dataset.update_dataset import ARDUpdater - if os.path.exists('src/.env'): from dotenv import load_dotenv load_dotenv() else: raise Exception("'src/.env' not found. Rename the 'src/.env.example' file and fill in values.") -OPENAI_API_KEY = os.environ.get('OPENAI_API_KEY') -openai.api_key = OPENAI_API_KEY +import openai +openai.api_key = os.environ.get('OPENAI_API_KEY') + +from dataset.update_dataset import ARDUpdater def update_database_and_pinecone(): @@ -20,6 +18,7 @@ def update_database_and_pinecone(): min_tokens_per_block=200, max_tokens_per_block=300, ) + updater.reset_dbs() updater.update(['gwern_blog'])