itismouad/pythonic-raqa-langchain-pinecone
0
1import os2from typing import List3 4import pinecone5from tqdm.auto import tqdm6from uuid import uuid47import arxiv8 9from langchain.document_loaders import PyPDFLoader10from langchain.text_splitter import RecursiveCharacterTextSplitter11from langchain.embeddings.openai import OpenAIEmbeddings12from langchain.embeddings import CacheBackedEmbeddings13from langchain.storage import LocalFileStore14from langchain.vectorstores import Pinecone15 16INDEX_BATCH_LIMIT = 10017 18class CharacterTextSplitter:19 def __init__(20 self,21 chunk_size: int = 1000,22 chunk_overlap: int = 200,23 ):24 assert (25 chunk_size > chunk_overlap26 ), "Chunk size must be greater than chunk overlap"27 28 self.chunk_size = chunk_size29 self.chunk_overlap = chunk_overlap30 31 self.text_splitter = RecursiveCharacterTextSplitter(32 chunk_size = self.chunk_size, # the character length of the chunk33 chunk_overlap = self.chunk_overlap, # the character length of the overlap between chunks34 length_function = len, # the length function - in this case, character length (aka the python len() fn.)35 36 )37 38 def split(self, text: str) -> List[str]:39 return self.text_splitter.split_text(text)40 41class ArxivLoader:42 43 def __init__(self, query : str = "Nuclear Fission", max_results : int = 5, encoding: str = "utf-8"):44 """"""45 self.query = query46 self.max_results = max_results47 48 self.paper_urls = []49 self.documents = []50 self.splitter = CharacterTextSplitter()51 52 def retrieve_urls(self):53 """"""54 arxiv_client = arxiv.Client()55 search = arxiv.Search(56 query = self.query,57 max_results = self.max_results,58 sort_by = arxiv.SortCriterion.Relevance59 )60 61 for result in arxiv_client.results(search):62 self.paper_urls.append(result.pdf_url)63 64 def load_documents(self):65 """"""66 for paper_url in self.paper_urls:67 loader = PyPDFLoader(paper_url)68 69 self.documents.append(loader.load())70 71 def format_document(self, document):72 """"""73 metadata = {74 'source_document' : document.metadata["source"],75 'page_number' : document.metadata["page"]76 }77 78 record_texts = self.splitter.split(document.page_content)79 record_metadatas = [{80 "chunk": j, "text": text, **metadata81 } for j, text in enumerate(record_texts)]82 83 return record_texts, record_metadatas84 85 def main(self):86 """"""87 self.retrieve_urls()88 self.load_documents()89 90 91class PineconeIndexer:92 93 def __init__(self, index_name : str = "arxiv-paper-index", metric : str = "cosine", n_dims : int = 1536):94 """"""95 pinecone.init(96 api_key=os.environ["PINECONE_API_KEY"],97 environment=os.environ["PINECONE_ENV"]98 )99 100 if index_name not in pinecone.list_indexes():101 # we create a new index102 pinecone.create_index(103 name=index_name,104 metric=metric,105 dimension=n_dims106 )107 108 self.arxiv_loader = ArxivLoader()109 110 self.index = pinecone.Index(index_name)111 112 def load_embedder(self):113 """"""114 store = LocalFileStore("./cache/")115 116 core_embeddings_model = OpenAIEmbeddings()117 118 self.embedder = CacheBackedEmbeddings.from_bytes_store(119 core_embeddings_model,120 store,121 namespace=core_embeddings_model.model122 )123 124 def upsert(self, texts, metadatas):125 """"""126 ids = [str(uuid4()) for _ in range(len(texts))]127 embeds = self.embedder.embed_documents(texts)128 self.index.upsert(vectors=zip(ids, embeds, metadatas))129 130 def index_documents(self, documents, batch_limit : int = INDEX_BATCH_LIMIT):131 """"""132 texts = []133 metadatas = []134 135 # iterate through your top-level document136 for i in tqdm(range(len(documents))):137 138 # select single document object139 for page in documents[i] : 140 141 record_texts, record_metadatas = self.arxiv_loader.format_document(page)142 143 texts.extend(record_texts)144 metadatas.extend(record_metadatas)145 146 if len(texts) >= batch_limit:147 self.upsert(texts, metadatas)148 149 texts = []150 metadatas = []151 152 if len(texts) > 0:153 self.upsert(texts, metadatas)154 155 def get_vectorstore(self):156 """"""157 return Pinecone(self.index, self.embedder.embed_query, "text")158 159 160if __name__ == "__main__":161 162 print("-------------- Loading Arxiv --------------")163 axloader = ArxivLoader()164 axloader.retrieve_urls()165 axloader.load_documents()166 167 print("\n-------------- Splitting sample doc --------------")168 sample_doc = axloader.documents[0]169 sample_page = sample_doc[0]170 171 splitter = CharacterTextSplitter()172 chunks = splitter.split(sample_page.page_content)173 print(len(chunks))174 print(chunks[0])175 176 print("\n-------------- testing pinecode indexer --------------")177 178 pi = PineconeIndexer()179 pi.load_embedder()180 pi.index_documents(axloader.documents)181 182 print(pi.index.describe_index_stats())183 