mirror of
https://github.com/wassname/stampy-chat.git
synced 2026-09-12 13:00:42 +08:00
clean up
This commit is contained in:
+1
-6
@@ -142,13 +142,8 @@ dmypy.json
|
||||
temp/
|
||||
*tmp.py
|
||||
|
||||
api/dataset.pkl
|
||||
api/dataset_big.pkl
|
||||
api/dataset_300.pkl
|
||||
|
||||
api/.env.backup
|
||||
|
||||
src/dataset_tests.ipynb
|
||||
src/ARD_LangChain_QA_Chat.ipynb
|
||||
|
||||
src/dataset/data/*
|
||||
src/dataset/data/ARD.db
|
||||
@@ -13,19 +13,24 @@ ARD_DATASET_NAME = "StampyAI/alignment-research-dataset"
|
||||
|
||||
### EMBEDDINGS ###
|
||||
USE_OPENAI_EMBEDDINGS = False
|
||||
|
||||
OPENAI_EMBEDDINGS_MODEL = "text-embedding-ada-002"
|
||||
EMBEDDINGS_DIMS = 1536
|
||||
OPENAI_EMBEDDINGS_DIMS = 1536
|
||||
OPENAI_EMBEDDINGS_RATE_LIMIT = 3500
|
||||
|
||||
SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL = "sentence-transformers/multi-qa-mpnet-base-cos-v1"
|
||||
SENTENCE_TRANSFORMER_EMBEDDINGS_DIMS = 768
|
||||
|
||||
DEVICE = "cuda" if torch.cuda.is_available() else "cpu"
|
||||
|
||||
### PINECONE ###
|
||||
PINECONE_INDEX_NAME = "stampy-chat-embeddings-test"
|
||||
PINECONE_VALUES_DIMS = EMBEDDINGS_DIMS
|
||||
PINECONE_VALUES_DIMS = OPENAI_EMBEDDINGS_DIMS if USE_OPENAI_EMBEDDINGS else SENTENCE_TRANSFORMER_EMBEDDINGS_DIMS
|
||||
PINECONE_METRIC = "cosine"
|
||||
PINECONE_METADATA_ENTRIES = ["entry_id", "source", "title", "authors", "text"]
|
||||
PINECONE_API_KEY = os.environ["PINECONE_API_KEY"]
|
||||
PINECONE_ENVIRONMENT = os.environ["PINECONE_ENVIRONMENT"]
|
||||
|
||||
### MISCELLANEOUS ###
|
||||
CHUNK_SIZE = 5000
|
||||
MAX_NUM_AUTHORS_IN_SIGNATURE = 3
|
||||
@@ -11,7 +11,7 @@ from .text_splitter import TokenSplitter
|
||||
from .sql_db_handler import SQLDB
|
||||
from .pinecone_db_handler import PineconeDB
|
||||
|
||||
from .settings import USE_OPENAI_EMBEDDINGS, OPENAI_EMBEDDINGS_MODEL, SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL, EMBEDDINGS_DIMS, OPENAI_EMBEDDINGS_RATE_LIMIT, DEVICE, ARD_DATASET_NAME, MAX_NUM_AUTHORS_IN_SIGNATURE
|
||||
from .settings import USE_OPENAI_EMBEDDINGS, OPENAI_EMBEDDINGS_MODEL, OPENAI_EMBEDDINGS_DIMS, OPENAI_EMBEDDINGS_RATE_LIMIT, SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL, SENTENCE_TRANSFORMER_EMBEDDINGS_DIMS, DEVICE, ARD_DATASET_NAME, CHUNK_SIZE, MAX_NUM_AUTHORS_IN_SIGNATURE
|
||||
|
||||
import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -32,14 +32,14 @@ class ARDUpdater:
|
||||
self.hf_embeddings = HuggingFaceEmbeddings(
|
||||
model_name=SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL,
|
||||
model_kwargs={'device': DEVICE},
|
||||
encode_kwargs={'show_progress_bar': True}
|
||||
encode_kwargs={'show_progress_bar': False}
|
||||
)
|
||||
|
||||
def update(self, custom_sources: List[str] = ['all']):
|
||||
for source in custom_sources:
|
||||
self.update_source(source)
|
||||
|
||||
def update_source(self, source: str, chunk_size: int = 100):
|
||||
def update_source(self, source: str):
|
||||
logger.info(f"Updating {source} entries...")
|
||||
|
||||
streamed_dataset = load_dataset(
|
||||
@@ -50,7 +50,7 @@ class ARDUpdater:
|
||||
self.is_sql_entry_upserted
|
||||
)
|
||||
|
||||
for batch in self.batchify(streamed_dataset, chunk_size):
|
||||
for batch in self.batchify(streamed_dataset):
|
||||
entries_batch = batch['entries_batch']
|
||||
chunks_batch = batch['chunks_batch']
|
||||
chunks_ids_batch = batch['chunks_ids_batch']
|
||||
@@ -71,7 +71,7 @@ class ARDUpdater:
|
||||
|
||||
logger.info(f"Successfully updated {source} entries.")
|
||||
|
||||
def batchify(self, iterable, chunk_size):
|
||||
def batchify(self, iterable):
|
||||
entries_batch = []
|
||||
chunks_batch = []
|
||||
chunks_ids_batch = []
|
||||
@@ -80,13 +80,13 @@ class ARDUpdater:
|
||||
chunks = self.token_splitter.split(entry['text'], f"Title: {entry['title']}, Authors: {get_authors_str(entry['authors'])}")
|
||||
chunks_ids = [f"{entry['id']}_{str(i).zfill(6)}" for i in range(len(chunks))]
|
||||
|
||||
# Add this entry's chunks to the current batch, even if it causes the batch size to exceed chunk_size.
|
||||
# Add this entry's chunks to the current batch, even if it causes the batch size to exceed CHUNK_SIZE.
|
||||
entries_batch.append(entry)
|
||||
chunks_batch.extend(chunks)
|
||||
chunks_ids_batch.extend(chunks_ids)
|
||||
|
||||
# If this batch is large enough, yield it and start a new one.
|
||||
if len(chunks_batch) >= chunk_size:
|
||||
if len(chunks_batch) >= CHUNK_SIZE:
|
||||
yield {'entries_batch': entries_batch, 'chunks_batch': chunks_batch, 'chunks_ids_batch': chunks_ids_batch}
|
||||
|
||||
entries_batch = []
|
||||
@@ -144,7 +144,7 @@ class ARDUpdater:
|
||||
|
||||
@retry(stop=stop_after_attempt(3))
|
||||
def get_openai_embeddings(self, chunks):
|
||||
embeddings = np.zeros((len(chunks), EMBEDDINGS_DIMS))
|
||||
embeddings = np.zeros((len(chunks), OPENAI_EMBEDDINGS_DIMS))
|
||||
rate_limit = OPENAI_EMBEDDINGS_RATE_LIMIT # TODO: use this rate_limit
|
||||
|
||||
openai_output = openai.Embedding.create(
|
||||
|
||||
Reference in New Issue
Block a user