CoolFace
Apppublic

xnetba/ChatPDF

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
backend.py147 linesDownload Raw Back to root
1import os2 3from langchain import FAISS, OpenAI, HuggingFaceHub, Cohere, PromptTemplate4from langchain.chains import RetrievalQA, ConversationalRetrievalChain5from langchain.embeddings import OpenAIEmbeddings, HuggingFaceEmbeddings, CohereEmbeddings6from langchain.memory import ConversationBufferMemory7from langchain.schema import Document8from langchain.text_splitter import RecursiveCharacterTextSplitter, CharacterTextSplitter, NLTKTextSplitter, \9    SpacyTextSplitter10from langchain.vectorstores import Chroma, ElasticVectorSearch11from pypdf import PdfReader12 13from schema import EmbeddingTypes, IndexerType, TransformType, BotType14 15 16class QnASystem:17 18    def read_and_load_pdf(self, f_data):19        pdf_data = PdfReader(f_data)20        documents = []21        for idx, page in enumerate(pdf_data.pages):22            documents.append(Document(page_content=page.extract_text(),23                                      metadata={"page_no": idx, "source": f_data.name}))24 25        self.documents = documents26 27    def document_transformer(self, transform_type: TransformType):28        match transform_type:29            case TransformType.CharacterTransform:30                t_type = CharacterTextSplitter(chunk_size=1000, chunk_overlap=20)31            case TransformType.RecursiveTransform:32                t_type = RecursiveCharacterTextSplitter(chunk_size=1000, chunk_overlap=20)33            case TransformType.NLTKTransform:34                t_type = NLTKTextSplitter()35            case TransformType.SpacyTransform:36                t_type = SpacyTextSplitter()37 38            case _:39                raise IndexError("Invalid Transformer Type")40 41        self.transformed_documents = t_type.split_documents(documents=self.documents)42 43    def generate_embeddings(self, embedding_type: EmbeddingTypes = EmbeddingTypes.OPENAI,44                            indexer_type: IndexerType = IndexerType.FAISS, **kwargs):45        temperature = kwargs.get("temperature", 0)46        max_tokens = kwargs.get("max_tokens", 512)47        match embedding_type:48            case EmbeddingTypes.OPENAI:49                os.environ["OPENAI_API_KEY"] = kwargs.get("api_key") or os.getenv("OPENAI_API_KEY")50                embeddings = OpenAIEmbeddings()51                llm = OpenAI(temperature=temperature, max_tokens=max_tokens)52            case EmbeddingTypes.HUGGING_FACE:53                embeddings = HuggingFaceEmbeddings(model_name=kwargs.get("model_name"))54                llm = HuggingFaceHub(repo_id=kwargs.get("model_name"),55                                     model_kwargs={"temperature": temperature, "max_tokens": max_tokens})56            case EmbeddingTypes.COHERE:57                embeddings = CohereEmbeddings(model=kwargs.get("model_name"), cohere_api_key=kwargs.get("api_key"))58                llm = Cohere(model=kwargs.get("model_name"), cohere_api_key=kwargs.get("api_key"),59                             model_kwargs={"temperature": temperature,60                                           "max_tokens": max_tokens})61            case _:62                raise IndexError("Invalid Embedding Type")63 64        match indexer_type:65            case IndexerType.FAISS:66                indexer = FAISS67            case IndexerType.CHROMA:68                indexer = Chroma()69 70            case IndexerType.ELASTICSEARCH:71                indexer = ElasticVectorSearch(elasticsearch_url=kwargs.get("elasticsearch_url"))72            case _:73                raise IndexError("Invalid Indexer Function")74 75        self.llm = llm76        self.indexer = indexer77        self.vector_store = indexer.from_documents(documents=self.transformed_documents, embedding=embeddings)78 79    def get_retriever(self, search_type="similarity", top_k=5, **kwargs):80        retriever = self.vector_store.as_retriever(search_type=search_type, search_kwargs={"k": top_k})81        self.retriever = retriever82 83    def get_prompt(self, bot_type: BotType, **kwargs):84        match bot_type:85            case BotType.qna:86                prompt = """87                You are a smart and helpful AI assistant, who answer the question given context88                {context}89                Question: {question}90                """91            case BotType.conversational:92                prompt = """93                Given the following conversation and a follow up question, 94                rephrase the follow up question to be a standalone question, in its original language.95                \nChat History:\n{chat_history}\nFollow Up Input: {question}\nStandalone question:96                """97        return PromptTemplate(input_variables=["context", "question", "chat_history"], template=prompt)98 99    def build_qa(self, qa_type: BotType, chain_type="stuff",100                 return_documents: bool = True, **kwargs):101        match qa_type:102            case BotType.qna:103                self.chain = RetrievalQA.from_chain_type(llm=self.llm, retriever=self.retriever, chain_type=chain_type,104                                                         return_source_documents=return_documents, verbose=True)105 106            case BotType.conversational:107                self.memory = ConversationBufferMemory(memory_key="chat_history", return_messages=True,108                                                       output_key="answer")109                self.chain = ConversationalRetrievalChain.from_llm(llm=self.llm, retriever=self.retriever,110                                                                   chain_type=chain_type,111                                                                   return_source_documents=return_documents,112                                                                   memory=self.memory, verbose=True)113 114            case _:115                raise IndexError("Invalid QA Type")116 117    def ask_question(self, query):118        if type(self.chain) == RetrievalQA:119            data = {"query": query}120        else:121            data = {"question": query}122        return self.chain(data)123 124    def build_chain(self, transform_type, embedding_type, indexer_type, **kwargs):125        if hasattr(self, "llm"):126            return self.chain127        self.document_transformer(transform_type)128        self.generate_embeddings(embedding_type=embedding_type,129                                 indexer_type=indexer_type, **kwargs)130        self.get_retriever(**kwargs)131        qa = self.build_qa(qa_type=kwargs.get("bot_type"), **kwargs)132        return qa133 134 135if __name__ == "__main__":136    qna = QnASystem()137    with open("../docs/Doc A.pdf", "rb") as f:138        qna.read_and_load_pdf(f)139        chain = qna.build_chain(140            transform_type=TransformType.RecursiveTransform,141            embedding_type=EmbeddingTypes.OPENAI, indexer_type=IndexerType.FAISS,142            chain_type="map_reduce", bot_type=BotType.conversational, return_documents=True143        )144        question = qna.ask_question(query="Hi! Summarize the document.")145        question = qna.ask_question(query="What happened from June 1984 to September 1996")146        print(question)147