mirror of
https://github.com/wassname/stampy-chat.git
synced 2026-09-10 12:40:44 +08:00
comments, refactor, new text_splitter, +bias
This commit is contained in:
@@ -3,15 +3,18 @@
|
||||
from typing import Dict, List, Union
|
||||
import numpy as np
|
||||
from tenacity import retry, stop_after_attempt
|
||||
from tqdm.auto import tqdm
|
||||
from datasets import load_dataset
|
||||
import openai
|
||||
|
||||
from .text_splitter import TokenSplitter
|
||||
from .text_splitter import ParagraphSentenceUnitTextSplitter
|
||||
from .sql_db_handler import SQLDB
|
||||
from .pinecone_db_handler import PineconeDB
|
||||
|
||||
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
|
||||
from .settings import USE_OPENAI_EMBEDDINGS, OPENAI_EMBEDDINGS_MODEL, \
|
||||
OPENAI_EMBEDDINGS_DIMS, OPENAI_EMBEDDINGS_RATE_LIMIT, \
|
||||
SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL, SENTENCE_TRANSFORMER_EMBEDDINGS_DIMS, \
|
||||
ARD_DATASET_NAME, CHUNK_SIZE, MAX_NUM_AUTHORS_IN_SIGNATURE, \
|
||||
EMBEDDING_LENGTH_BIAS
|
||||
|
||||
import logging
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -20,31 +23,50 @@ logger = logging.getLogger(__name__)
|
||||
class ARDUpdater:
|
||||
def __init__(
|
||||
self,
|
||||
min_tokens_per_block: int = 200, # Minimum number of tokens per block.
|
||||
max_tokens_per_block: int = 400, # Maximum number of tokens per block.
|
||||
min_chunk_size: int = ParagraphSentenceUnitTextSplitter.DEFAULT_MIN_CHUNK_SIZE,
|
||||
max_chunk_size: int = ParagraphSentenceUnitTextSplitter.DEFAULT_MAX_CHUNK_SIZE,
|
||||
):
|
||||
self.token_splitter = TokenSplitter(min_tokens_per_block, max_tokens_per_block)
|
||||
self.text_splitter = ParagraphSentenceUnitTextSplitter(
|
||||
min_chunk_size=min_chunk_size,
|
||||
max_chunk_size=max_chunk_size,
|
||||
)
|
||||
self.sql_db = SQLDB()
|
||||
self.pinecone_db = PineconeDB()
|
||||
|
||||
if not USE_OPENAI_EMBEDDINGS:
|
||||
import torch
|
||||
from langchain.embeddings import HuggingFaceEmbeddings
|
||||
|
||||
self.hf_embeddings = HuggingFaceEmbeddings(
|
||||
model_name=SENTENCE_TRANSFORMER_EMBEDDINGS_MODEL,
|
||||
model_kwargs={'device': DEVICE},
|
||||
model_kwargs={'device': "cuda" if torch.cuda.is_available() else "cpu"},
|
||||
encode_kwargs={'show_progress_bar': False}
|
||||
)
|
||||
|
||||
def update(self, custom_sources: List[str] = ['all']):
|
||||
"""
|
||||
Update the given sources. If no sources are provided, updates all sources.
|
||||
|
||||
:param custom_sources: List of sources to update.
|
||||
"""
|
||||
|
||||
for source in custom_sources:
|
||||
self.update_source(source)
|
||||
|
||||
def update_source(self, source: str):
|
||||
"""
|
||||
Updates the entries from the given source.
|
||||
|
||||
:param source: The name of the source to update.
|
||||
"""
|
||||
|
||||
logger.info(f"Updating {source} entries...")
|
||||
|
||||
streamed_dataset = load_dataset(
|
||||
ARD_DATASET_NAME, source, split='train', streaming=True
|
||||
).map(self.preprocess_and_validate).filter(
|
||||
).map(
|
||||
self.preprocess_and_validate
|
||||
).filter(
|
||||
self.is_valid_entry
|
||||
).filter(
|
||||
self.is_sql_entry_upserted
|
||||
@@ -54,53 +76,85 @@ class ARDUpdater:
|
||||
entries_batch = batch['entries_batch']
|
||||
chunks_batch = batch['chunks_batch']
|
||||
chunks_ids_batch = batch['chunks_ids_batch']
|
||||
sources_batch = batch['sources_batch']
|
||||
|
||||
try:
|
||||
if USE_OPENAI_EMBEDDINGS:
|
||||
embeddings = self.get_openai_embeddings(chunks_batch)
|
||||
else:
|
||||
embeddings = np.array(self.hf_embeddings.embed_documents(chunks_batch))
|
||||
embeddings = self.extract_embeddings(chunks_batch, sources_batch)
|
||||
|
||||
self.sql_db.upsert_chunks(chunks_ids_batch, chunks_batch)
|
||||
self.sql_db.upsert_chunks(chunks_ids_batch, chunks_batch, embeddings)
|
||||
self.pinecone_db.delete_entries([entry['id'] for entry in entries_batch])
|
||||
self.pinecone_db.upsert_entries(entries_batch, chunks_batch, chunks_ids_batch, embeddings)
|
||||
|
||||
logger.info(f"Successfully updated {len(entries_batch)} {source} entries with {len(chunks_batch)} total chunks.")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"An error occurred while updating source {source}: {str(e)}", exc_info=True)
|
||||
|
||||
logger.info(f"Successfully updated {source} entries.")
|
||||
|
||||
def batchify(self, iterable):
|
||||
"""
|
||||
Divides the iterable into batches of size ~CHUNK_SIZE.
|
||||
|
||||
:param iterable: The iterable to divide into batches.
|
||||
:returns: A generator that yields batches from the iterable.
|
||||
"""
|
||||
|
||||
entries_batch = []
|
||||
chunks_batch = []
|
||||
chunks_ids_batch = []
|
||||
sources_batch = []
|
||||
|
||||
for entry in iterable:
|
||||
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))]
|
||||
chunks, chunks_ids = self.create_chunk_ids_and_authors(entry)
|
||||
|
||||
# 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)
|
||||
sources_batch.extend([entry['source']] * len(chunks))
|
||||
|
||||
# If this batch is large enough, yield it and start a new one.
|
||||
if len(chunks_batch) >= CHUNK_SIZE:
|
||||
yield {'entries_batch': entries_batch, 'chunks_batch': chunks_batch, 'chunks_ids_batch': chunks_ids_batch}
|
||||
yield self._create_batch(entries_batch, chunks_batch, chunks_ids_batch, sources_batch)
|
||||
|
||||
entries_batch = []
|
||||
chunks_batch = []
|
||||
chunks_ids_batch = []
|
||||
sources_batch = []
|
||||
|
||||
# Yield any remaining items.
|
||||
if entries_batch:
|
||||
yield {'entries_batch': entries_batch, 'chunks_batch': chunks_batch, 'chunks_ids_batch': chunks_ids_batch}
|
||||
|
||||
def preprocess_and_validate(self, entry):
|
||||
yield self._create_batch(entries_batch, chunks_batch, chunks_ids_batch, sources_batch)
|
||||
|
||||
def create_chunk_ids_and_authors(self, entry):
|
||||
signature = f"Title: {entry['title']}, Author(s): {self.get_authors_str(entry['authors'])}"
|
||||
chunks = self.text_splitter.split_text(entry['text'])
|
||||
chunks = [f"- {signature}\n\n{chunk}" for chunk in chunks]
|
||||
chunks_ids = [f"{entry['id']}_{str(i).zfill(6)}" for i in range(len(chunks))]
|
||||
return chunks, chunks_ids
|
||||
|
||||
def _create_batch(self, entries_batch, chunks_batch, chunks_ids_batch, sources_batch):
|
||||
return {'entries_batch': entries_batch, 'chunks_batch': chunks_batch, 'chunks_ids_batch': chunks_ids_batch, 'sources_batch': sources_batch}
|
||||
|
||||
def is_sql_entry_upserted(self, entry):
|
||||
"""Upserts an entry to the SQL database and returns the success status"""
|
||||
return self.sql_db.upsert_entry(entry)
|
||||
|
||||
def extract_embeddings(self, chunks_batch, sources_batch):
|
||||
if USE_OPENAI_EMBEDDINGS:
|
||||
return self.get_openai_embeddings(chunks_batch, sources_batch)
|
||||
else:
|
||||
return np.array(self.hf_embeddings.embed_documents(chunks_batch, sources_batch))
|
||||
|
||||
def reset_dbs(self):
|
||||
self.sql_db.create_tables(True)
|
||||
self.pinecone_db.create_index(True)
|
||||
|
||||
@staticmethod
|
||||
def preprocess_and_validate(entry):
|
||||
"""Preprocesses and validates the entry data"""
|
||||
try:
|
||||
self.validate_entry(entry)
|
||||
ARDUpdater.validate_entry(entry)
|
||||
|
||||
return {
|
||||
'id': entry['id'],
|
||||
@@ -115,7 +169,8 @@ class ARDUpdater:
|
||||
logger.error(f"Entry validation failed: {str(e)}", exc_info=True)
|
||||
return None
|
||||
|
||||
def validate_entry(self, entry: Dict[str, Union[str, list]], char_len_lower_limit: int = 0):
|
||||
@staticmethod
|
||||
def validate_entry(entry: Dict[str, Union[str, list]], char_len_lower_limit: int = 0):
|
||||
metadata_types = {
|
||||
'id': str,
|
||||
'source': str,
|
||||
@@ -138,36 +193,26 @@ class ARDUpdater:
|
||||
"""Checks if the entry is valid"""
|
||||
return entry is not None
|
||||
|
||||
def is_sql_entry_upserted(self, entry):
|
||||
"""Upserts an entry to the SQL database and returns the success status"""
|
||||
return self.sql_db.upsert_entry(entry)
|
||||
|
||||
@retry(stop=stop_after_attempt(3))
|
||||
def get_openai_embeddings(self, chunks):
|
||||
@staticmethod
|
||||
def get_openai_embeddings(chunks, sources=''):
|
||||
embeddings = np.zeros((len(chunks), OPENAI_EMBEDDINGS_DIMS))
|
||||
rate_limit = OPENAI_EMBEDDINGS_RATE_LIMIT # TODO: use this rate_limit
|
||||
|
||||
openai_output = openai.Embedding.create(
|
||||
model=OPENAI_EMBEDDINGS_MODEL,
|
||||
input=chunks
|
||||
)['data']
|
||||
|
||||
for i, embedding in enumerate(openai_output):
|
||||
embeddings[i] = embedding['embedding']
|
||||
for i, (embedding, source) in enumerate(zip(openai_output, sources)):
|
||||
bias = EMBEDDING_LENGTH_BIAS.get(source, 1.0)
|
||||
embeddings[i] = bias * np.array(embedding['embedding'])
|
||||
|
||||
return embeddings
|
||||
|
||||
def reset_dbs(self):
|
||||
self.sql_db.create_tables(True)
|
||||
self.pinecone_db.create_index(True)
|
||||
|
||||
|
||||
##### Helper functions #####
|
||||
|
||||
def get_authors_str(authors_lst: List[str]) -> str:
|
||||
if authors_lst == []: return 'n/a'
|
||||
if len(authors_lst) == 1: return authors_lst[0]
|
||||
else:
|
||||
authors_lst = authors_lst[:MAX_NUM_AUTHORS_IN_SIGNATURE]
|
||||
authors_str = f"{', '.join(authors_lst[:-1])} and {authors_lst[-1]}"
|
||||
return authors_str
|
||||
@staticmethod
|
||||
def get_authors_str(authors_lst: List[str]) -> str:
|
||||
if authors_lst == []: return 'n/a'
|
||||
if len(authors_lst) == 1: return authors_lst[0]
|
||||
else:
|
||||
authors_lst = authors_lst[:MAX_NUM_AUTHORS_IN_SIGNATURE]
|
||||
authors_str = f"{', '.join(authors_lst[:-1])} and {authors_lst[-1]}"
|
||||
return authors_str
|
||||
Reference in New Issue
Block a user