xnetba/ChatPDF
0
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 