jamesredd/document_analysis
0
1import openai2import gradio as gr3import os4 5from langchain.document_loaders import UnstructuredFileLoader 6 7from langchain.embeddings.openai import OpenAIEmbeddings8from langchain.vectorstores import Chroma9 10from langchain.chains import RetrievalQA11from langchain.chat_models import ChatOpenAI12from langchain.text_splitter import RecursiveCharacterTextSplitter13 14class DocumentManager:15 def __init__(self):16 self.api_key = None17 self.citation = ""18 self.docs = []19 self.retriever = None20 self.files = []21 self.provide_citation = True 22 23 self.source_documents = []24 25 self.user_prompt = "Be direct and cite your sources."26 self.text_splitter = RecursiveCharacterTextSplitter(chunk_size=2000, chunk_overlap=100)27 28 self.check_for_api_key_env()29 30 def check_for_api_key_env(self):31 if "OPENAI_API_KEY" in os.environ:32 self.set_api_key(os.environ["OPENAI_API_KEY"])33 34 def set_api_key(self, value):35 if (value is None) or (not value.startswith("sk")):36 gr.Warning("Please enter a valid OpenAI API key.")37 return self.create_api_key_status_display()38 39 self.api_key = value40 openai.api_key = self.api_key41 if len(self.docs) > 0:42 documents = self.text_splitter.split_documents(self.docs)43 self.retriever = Chroma.from_documents(documents, OpenAIEmbeddings(openai_api_key=self.api_key)).as_retriever(search_type="mmr", search_kwargs={'fetch_k': 30}, return_source_documents=True)44 else:45 self.retriever = Chroma(embedding_function=OpenAIEmbeddings(openai_api_key=self.api_key)).as_retriever(search_type="mmr", search_kwargs={'fetch_k': 30}, return_source_documents=True)46 self.llm = ChatOpenAI(model_name="gpt-4", temperature=0, streaming=True, openai_api_key=self.api_key)47 self.qa = RetrievalQA.from_chain_type(48 llm=self.llm,49 chain_type="stuff", 50 retriever=self.retriever,51 return_source_documents=True)52 53 return self.create_api_key_status_display()54 55 def create_api_key_status_display(self):56 if self.api_key is None:57 return gr.Textbox("❌ Please enter an API key.", label=None, interactive=False, container=False)58 else:59 return gr.Textbox(f"✅ API key: {self.api_key[:9]}", label=None, interactive=False, container=False)60 61 def get_user_prompt(self): 62 return self.user_prompt63 64 def set_user_prompt(self, value):65 self.user_prompt = value66 67 def set_provide_citation(self, value):68 self.provide_citation = value69 70 def delete_files(self):71 self.docs = []72 self.files = []73 self.source_documents = []74 self.db = Chroma(embedding_function=OpenAIEmbeddings(openai_api_key=self.api_key))75 76 self.db._client_settings.allow_reset = True77 self.db._client.reset()78 79 self.retreiver = self.db.as_retriever(search_type="mmr", search_kwargs={'fetch_k': 30}, return_source_documents=True)80 self.llm = ChatOpenAI(model_name="gpt-4", temperature=0, streaming=True, openai_api_key=self.api_key)81 82 self.qa = RetrievalQA.from_chain_type(83 llm=self.llm,84 chain_type="stuff", 85 retriever=Chroma(86 embedding_function=OpenAIEmbeddings(openai_api_key=self.api_key))87 .as_retriever(search_type="mmr", search_kwargs={'fetch_k': 30}, return_source_documents=True),88 return_source_documents=True)89 90 return gr.Markdown(self.generate_file_markdown(), label="Uploaded files")91 92 def reset_qa(self):93 documents = self.text_splitter.split_documents(self.docs)94 self.retriever = Chroma.from_documents(documents, OpenAIEmbeddings(openai_api_key=self.api_key)).as_retriever(search_type="mmr", search_kwargs={'fetch_k': 30}, return_source_documents=True)95 self.llm = ChatOpenAI(model_name="gpt-4", temperature=0, streaming=True, openai_api_key=self.api_key)96 self.qa = RetrievalQA.from_chain_type(97 llm=self.llm,98 chain_type="stuff", 99 retriever=self.retriever,100 return_source_documents=True)101 102 def tokenize_doc(self, filepath):103 loader = UnstructuredFileLoader(filepath)104 doc = loader.load()105 self.docs.extend(doc)106 107 def update_citation(self):108 if self.provide_citation and self.api_key:109 summed = ""110 for doc in self.source_documents:111 summed += doc.page_content112 113 self.citation = self.llm.predict(114 "Question: " + self.question + ". Answer: " + self.result + ". Citation: " + summed + 115 ". From the citation, return the relevant passage and the exact articles")116 else:117 self.citation = ""118 119 return self.citation120 121 def predict(self, message, history):122 if self.api_key is None:123 gr.Warning("Please enter an OpenAI API key in the settings tab.")124 return "", []125 126 if history is None:127 history = []128 129 summed_history = " ".join(sum(history, []))130 question = "You have access to these documents:" + self.generate_file_markdown() + ". Do not make things up, only say what you have a primary source document for. ---- CHAT HISTORY : " + summed_history + " --- SYSTEM PROMPT: " + self.user_prompt + " -- Answer this question: " + message131 132 print(self.retriever.vectorstore.get())133 134 result = self.qa({"query": question})135 136 self.source_documents = result["source_documents"]137 self.result = result["result"]138 self.question = question139 140 history.append([message, ""])141 for message in result["result"]:142 history[-1][1] += message143 yield "", history144 145 def generate_file_markdown(self):146 files_md = ""147 for file in self.files:148 filename = file.split("/")[-1]149 files_md += "- " + filename + "\n"150 return files_md151 152 def upload_file(self, files):153 if self.api_key is None:154 gr.Warning("Please enter an OpenAI API key.")155 return self.files156 157 for file in files:158 self.tokenize_doc(file.name)159 filepaths = [file.orig_name for file in files]160 self.files = filepaths + self.files161 self.reset_qa()162 163 return gr.Markdown(self.generate_file_markdown(), label="Uploaded files") 164 165 def create_delete_button(self, value):166 if value and self.api_key:167 return gr.Button("Delete files", scale=4, interactive=True)168 else:169 return gr.Button("Delete files", scale=4, interactive=False)170 171def create_demo():172 doc_manager = DocumentManager()173 174 with gr.Blocks() as demo:175 with gr.Tab("Chat"):176 with gr.Row():177 chatbot = gr.Chatbot(scale=5, layout="panel", height=700)178 with gr.Column():179 citation = gr.Textbox("", label="Citation", interactive=False, scale=3, container=False)180 checkbox = gr.Checkbox(label="Provide document citation", value=True)181 checkbox.change(doc_manager.set_provide_citation, checkbox)182 msg = gr.Textbox(label="Enter your message")183 with gr.Row():184 submit_button = gr.Button("Submit ➡️")185 submit_button.click(doc_manager.predict, [msg, chatbot], [msg, chatbot]).then(doc_manager.update_citation, None, citation)186 clear = gr.ClearButton([msg, chatbot, citation])187 188 msg.submit(doc_manager.predict, [msg, chatbot], [msg, chatbot]).then(doc_manager.update_citation, None, citation)189 190 with gr.Tab("Settings") as settings_tab:191 with gr.Row():192 api_key_textbox = gr.Textbox(label="OpenAI API Key", scale=5)193 with gr.Column():194 api_key_status = doc_manager.create_api_key_status_display()195 save_key_button = gr.Button("Save Key")196 save_key_button.click(doc_manager.set_api_key, inputs=api_key_textbox, outputs=api_key_status).then(lambda:None, None, api_key_textbox, queue=False)197 198 file_output = gr.Markdown("", label="Uploaded files")199 upload_button = gr.UploadButton("Upload Files", file_count="multiple")200 upload_button.upload(doc_manager.upload_file, upload_button, file_output)201 prompt_textbox = gr.Textbox(label="Prompt", value=doc_manager.get_user_prompt())202 prompt_textbox.change(doc_manager.set_user_prompt, prompt_textbox)203 204 with gr.Row():205 allow_delete_checkbox = gr.Checkbox(value=False, label="Allow deletion of files") 206 delete_button = doc_manager.create_delete_button(False)207 delete_button.click(doc_manager.delete_files, outputs=file_output)208 allow_delete_checkbox.select(doc_manager.create_delete_button, outputs=delete_button, inputs=allow_delete_checkbox)209 210 settings_tab.select(doc_manager.create_api_key_status_display, outputs=api_key_status)211 212 return demo213 214 215if __name__ == "__main__":216 demo = create_demo() 217 demo.queue().launch(auth=("user", "pw"))218 