diff --git a/src/dataset/settings.py b/src/dataset/settings.py index f524191..951e138 100644 --- a/src/dataset/settings.py +++ b/src/dataset/settings.py @@ -1,6 +1,7 @@ # dataset/settings.py import os +import torch from pathlib import Path ### FILE PATHS ### @@ -13,9 +14,10 @@ ARD_DATASET_NAME = "StampyAI/alignment-research-dataset" ### EMBEDDINGS ### USE_OPENAI_EMBEDDINGS = False OPENAI_EMBEDDINGS_MODEL = "text-embedding-ada-002" -SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL = "sentence-transformers/multi-qa-mpnet-base-cos-v1" EMBEDDINGS_DIMS = 1536 OPENAI_EMBEDDINGS_RATE_LIMIT = 3500 +SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL = "sentence-transformers/multi-qa-mpnet-base-cos-v1" +DEVICE = "cuda" if torch.cuda.is_available() else "cpu" ### PINECONE ### PINECONE_INDEX_NAME = "stampy-chat-embeddings-test" diff --git a/src/dataset/update_dataset.py b/src/dataset/update_dataset.py index 671a8d4..279b20a 100644 --- a/src/dataset/update_dataset.py +++ b/src/dataset/update_dataset.py @@ -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, ARD_DATASET_NAME, MAX_NUM_AUTHORS_IN_SIGNATURE +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 import logging logger = logging.getLogger(__name__) @@ -31,6 +31,8 @@ class ARDUpdater: from langchain.embeddings import HuggingFaceEmbeddings self.hf_embeddings = HuggingFaceEmbeddings( model_name=SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL, + model_kwargs={'device': DEVICE}, + encode_kwargs={'show_progress_bar': True} ) def update(self, custom_sources: List[str] = ['all']):