CoolFace
Apppublic

jamesredd/document_analysis

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
app.py218 linesDownload Raw Back to root
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