CoolFace
Apppublic

Sambhavnoobcoder/Cyber-Security-Advisor

sourceHugging Faceupdated 2y agoView on Hugging Face
1likes
app.py217 linesDownload Raw Back to root
1import gradio as gr2import os3api_token = os.getenv("HF_TOKEN")4 5 6from langchain_community.vectorstores import FAISS7from langchain_community.document_loaders import PyPDFLoader8from langchain.text_splitter import RecursiveCharacterTextSplitter9from langchain_community.vectorstores import Chroma10from langchain.chains import ConversationalRetrievalChain11from langchain_community.embeddings import HuggingFaceEmbeddings 12from langchain_community.llms import HuggingFacePipeline13from langchain.chains import ConversationChain14from langchain.memory import ConversationBufferMemory15from langchain_community.llms import HuggingFaceEndpoint16import torch17 18list_llm = ["meta-llama/Meta-Llama-3-8B-Instruct", "mistralai/Mistral-7B-Instruct-v0.2"]  19list_llm_simple = [os.path.basename(llm) for llm in list_llm]20 21# Load and split PDF document22def load_doc(list_file_path):23    # Processing for one document only24    # loader = PyPDFLoader(file_path)25    # pages = loader.load()26    loaders = [PyPDFLoader(x) for x in list_file_path]27    pages = []28    for loader in loaders:29        pages.extend(loader.load())30    text_splitter = RecursiveCharacterTextSplitter(31        chunk_size = 1024, 32        chunk_overlap = 64 33    )  34    doc_splits = text_splitter.split_documents(pages)35    return doc_splits36 37# Create vector database38def create_db(splits):39    embeddings = HuggingFaceEmbeddings()40    vectordb = FAISS.from_documents(splits, embeddings)41    return vectordb42 43 44# Initialize langchain LLM chain45def initialize_llmchain(llm_model, temperature, max_tokens, top_k, vector_db, progress=gr.Progress()):46    if llm_model == "meta-llama/Meta-Llama-3-8B-Instruct":47        llm = HuggingFaceEndpoint(48            repo_id=llm_model,49            huggingfacehub_api_token = api_token,50            temperature = temperature,51            max_new_tokens = max_tokens,52            top_k = top_k,53        )54    else:55        llm = HuggingFaceEndpoint(56            huggingfacehub_api_token = api_token,57            repo_id=llm_model, 58            temperature = temperature,59            max_new_tokens = max_tokens,60            top_k = top_k,61        )62    63    memory = ConversationBufferMemory(64        memory_key="chat_history",65        output_key='answer',66        return_messages=True67    )68 69    retriever=vector_db.as_retriever()70    qa_chain = ConversationalRetrievalChain.from_llm(71        llm,72        retriever=retriever,73        chain_type="stuff", 74        memory=memory,75        return_source_documents=True,76        verbose=False,77    )78    return qa_chain79 80# Initialize database81def initialize_database(list_file_obj, progress=gr.Progress()):82    # Create a list of documents (when valid)83    list_file_path = [x.name for x in list_file_obj if x is not None]84    # Load document and create splits85    doc_splits = load_doc(list_file_path)86    # Create or load vector database87    vector_db = create_db(doc_splits)88    return vector_db, "Database created!"89 90# Initialize LLM91def initialize_LLM(llm_option, llm_temperature, max_tokens, top_k, vector_db, progress=gr.Progress()):92    # print("llm_option",llm_option)93    llm_name = list_llm[llm_option]94    print("llm_name: ",llm_name)95    qa_chain = initialize_llmchain(llm_name, llm_temperature, max_tokens, top_k, vector_db, progress)96    return qa_chain, "QA chain initialized. Chatbot is ready!"97 98 99def format_chat_history(message, chat_history):100    formatted_chat_history = []101    for user_message, bot_message in chat_history:102        formatted_chat_history.append(f"User: {user_message}")103        formatted_chat_history.append(f"Assistant: {bot_message}")104    return formatted_chat_history105    106 107def conversation(qa_chain, message, history):108    formatted_chat_history = format_chat_history(message, history)109    # Generate response using QA chain110    response = qa_chain.invoke({"question": message, "chat_history": formatted_chat_history})111    response_answer = response["answer"]112    if response_answer.find("Helpful Answer:") != -1:113        response_answer = response_answer.split("Helpful Answer:")[-1]114    response_sources = response["source_documents"]115    response_source1 = response_sources[0].page_content.strip()116    response_source2 = response_sources[1].page_content.strip()117    response_source3 = response_sources[2].page_content.strip()118    # Langchain sources are zero-based119    response_source1_page = response_sources[0].metadata["page"] + 1120    response_source2_page = response_sources[1].metadata["page"] + 1121    response_source3_page = response_sources[2].metadata["page"] + 1122    # Append user message and response to chat history123    new_history = history + [(message, response_answer)]124    return qa_chain, gr.update(value=""), new_history, response_source1, response_source1_page, response_source2, response_source2_page, response_source3, response_source3_page125    126 127def upload_file(file_obj):128    list_file_path = []129    for idx, file in enumerate(file_obj):130        file_path = file_obj.name131        list_file_path.append(file_path)132    return list_file_path133 134 135def demo():136    # with gr.Blocks(theme=gr.themes.Default(primary_hue="sky")) as demo:137    with gr.Blocks(theme=gr.themes.Default(primary_hue="red", secondary_hue="pink", neutral_hue = "sky")) as demo:138        vector_db = gr.State()139        qa_chain = gr.State()140        gr.HTML("<center><h1>RAG PDF chatbot</h1><center>")141        gr.Markdown("""<b>Query your PDF documents!</b> This AI agent is designed to perform retrieval augmented generation (RAG) on PDF documents. The app is hosted on Hugging Face Hub for the sole purpose of demonstration. \142        <b>Please do not upload confidential documents.</b>143        """)144        with gr.Row():145            with gr.Column(scale = 86):146                gr.Markdown("<b>Step 1 - Upload PDF documents and Initialize RAG pipeline</b>")147                with gr.Row():148                    document = gr.Files(height=300, file_count="multiple", file_types=["pdf"], interactive=True, label="Upload PDF documents")149                with gr.Row():150                    db_btn = gr.Button("Create vector database")151                with gr.Row():152                        db_progress = gr.Textbox(value="Not initialized", show_label=False) # label="Vector database status", 153                gr.Markdown("<style>body { font-size: 16px; }</style><b>Select Large Language Model (LLM) and input parameters</b>")154                with gr.Row():155                    llm_btn = gr.Radio(list_llm_simple, label="Available LLMs", value = list_llm_simple[0], type="index") # info="Select LLM", show_label=False156                with gr.Row():157                    with gr.Accordion("LLM input parameters", open=False):158                        with gr.Row():159                            slider_temperature = gr.Slider(minimum = 0.01, maximum = 1.0, value=0.5, step=0.1, label="Temperature", info="Controls randomness in token generation", interactive=True)160                        with gr.Row():161                            slider_maxtokens = gr.Slider(minimum = 128, maximum = 9192, value=4096, step=128, label="Max New Tokens", info="Maximum number of tokens to be generated",interactive=True)162                        with gr.Row():163                                slider_topk = gr.Slider(minimum = 1, maximum = 10, value=3, step=1, label="top-k", info="Number of tokens to select the next token from", interactive=True)164                with gr.Row():165                    qachain_btn = gr.Button("Initialize Question Answering Chatbot")166                with gr.Row():167                        llm_progress = gr.Textbox(value="Not initialized", show_label=False) # label="Chatbot status", 168 169            with gr.Column(scale = 200):170                gr.Markdown("<b>Step 2 - Chat with your Document</b>")171                chatbot = gr.Chatbot(height=505)172                with gr.Accordion("Relevent context from the source document", open=False):173                    with gr.Row():174                        doc_source1 = gr.Textbox(label="Reference 1", lines=2, container=True, scale=20)175                        source1_page = gr.Number(label="Page", scale=1)176                    with gr.Row():177                        doc_source2 = gr.Textbox(label="Reference 2", lines=2, container=True, scale=20)178                        source2_page = gr.Number(label="Page", scale=1)179                    with gr.Row():180                        doc_source3 = gr.Textbox(label="Reference 3", lines=2, container=True, scale=20)181                        source3_page = gr.Number(label="Page", scale=1)182                with gr.Row():183                    msg = gr.Textbox(placeholder="Ask a question", container=True)184                with gr.Row():185                    submit_btn = gr.Button("Submit")186                    clear_btn = gr.ClearButton([msg, chatbot], value="Clear")187            188        # Preprocessing events189        db_btn.click(initialize_database, \190            inputs=[document], \191            outputs=[vector_db, db_progress])192        qachain_btn.click(initialize_LLM, \193            inputs=[llm_btn, slider_temperature, slider_maxtokens, slider_topk, vector_db], \194            outputs=[qa_chain, llm_progress]).then(lambda:[None,"",0,"",0,"",0], \195            inputs=None, \196            outputs=[chatbot, doc_source1, source1_page, doc_source2, source2_page, doc_source3, source3_page], \197            queue=False)198 199        # Chatbot events200        msg.submit(conversation, \201            inputs=[qa_chain, msg, chatbot], \202            outputs=[qa_chain, msg, chatbot, doc_source1, source1_page, doc_source2, source2_page, doc_source3, source3_page], \203            queue=False)204        submit_btn.click(conversation, \205            inputs=[qa_chain, msg, chatbot], \206            outputs=[qa_chain, msg, chatbot, doc_source1, source1_page, doc_source2, source2_page, doc_source3, source3_page], \207            queue=False)208        clear_btn.click(lambda:[None,"",0,"",0,"",0], \209            inputs=None, \210            outputs=[chatbot, doc_source1, source1_page, doc_source2, source2_page, doc_source3, source3_page], \211            queue=False)212    demo.queue().launch(debug=True)213 214 215if __name__ == "__main__":216    demo()217