bug-fix and fixed device option

This commit is contained in:
henri123lemoine
2023-07-02 07:22:38 -04:00
parent cdfb5a52a7
commit aeba2fbfb0
2 changed files with 6 additions and 2 deletions
+3 -1
View File
@@ -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"
+3 -1
View File
@@ -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']):