Team Ai
Apppublic

itismouad/pythonic-raqa-langchain-pinecone

sourceHugging Faceopenrailupdated 3y agoView on Hugging Face
0likes
utils.py183 linesDownload Raw Back to root
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