mirror of
https://github.com/wassname/stampy-chat.git
synced 2026-09-20 13:20:52 +08:00
Renamed classes, added settings
This commit is contained in:
@@ -10,7 +10,7 @@ import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
class PineconeDB:
|
||||
class PineconeDBHandler:
|
||||
def __init__(
|
||||
self,
|
||||
create_index: bool = False,
|
||||
|
||||
+17
-2
@@ -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"]
|
||||
PINECONE_METADATA_ENTRIES = ["entry_id", "source", "title", "authors", "text"]
|
||||
|
||||
### MISC ###
|
||||
CUSTOM_SOURCES = ['gwern_blog']
|
||||
@@ -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()
|
||||
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
+5
-6
@@ -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'])
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user