From aeba2fbfb05c10f74295e92837081130b3040443 Mon Sep 17 00:00:00 2001 From: henri123lemoine Date: Sun, 2 Jul 2023 07:22:38 -0400 Subject: [PATCH] bug-fix and fixed device option --- src/dataset/settings.py | 4 +++- src/dataset/update_dataset.py | 4 +++- 2 files changed, 6 insertions(+), 2 deletions(-) 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']):