Team Ai
Apppublic

MathWizard1729/PDF_Chatbot_Gradio_ChromaDB

sourceHugging Faceupdated 1y agoView on Hugging Face
0likes
app.py235 linesDownload Raw Back to root
1import streamlit as st2import boto33from langchain_community.document_loaders import PyPDFLoader4from langchain_text_splitters import RecursiveCharacterTextSplitter5from langchain_aws import BedrockEmbeddings6from langchain_chroma import Chroma7from langchain_aws import ChatBedrock8from langchain.prompts import ChatPromptTemplate9from langchain.schema import StrOutputParser10from langchain.schema.runnable import RunnablePassthrough11import os12 13# --- Streamlit UI Setup (MUST BE THE FIRST STREAMLIT COMMAND) ---14st.set_page_config(15    page_title="Math Research Paper RAG Bot",16    page_icon="๐Ÿ“š",17    layout="wide"18)19 20st.title("๐Ÿ“š Math Research Paper RAG Chatbot")21st.markdown(22    """23    Upload a mathematical research paper (PDF) and ask questions about its content. 24    This bot uses Amazon Bedrock (Claude 3 Sonnet for reasoning, Titan Embeddings for vectors) 25    and ChromaDB for Retrieval-Augmented Generation.26    27    **Note:** This application requires AWS credentials (`AWS_ACCESS_KEY_ID`, `AWS_SECRET_ACCESS_KEY`) 28    and region (`AWS_REGION`) to be set up in a `.env` file or environment variables.29    """30)31 32# --- Configuration ---33# Set AWS region (adjust if needed, loaded from .env or env var)34AWS_REGION = os.getenv("AWS_REGION") 35if not AWS_REGION:36    st.error("AWS_REGION not found in environment variables or .env file. Please set it.")37    st.stop()38 39# Bedrock model IDs40EMBEDDING_MODEL_ID = "amazon.titan-embed-text-v1"41# Claude 4 is not generally available via Bedrock. Using Claude 3 Sonnet.42LLM_MODEL_ID = "anthropic.claude-3-sonnet-20240229-v1:0" 43 44# --- Initialize Bedrock Client (once) ---45@st.cache_resource46def get_bedrock_client():47    """Initializes and returns a boto3 Bedrock client.48    Returns: Tuple (boto3_client, success_bool, error_message_str or None)49    """50    try:51        client = boto3.client(52            service_name="bedrock-runtime",53            region_name=AWS_REGION54        )55        # Optional: Verify credentials by trying a simple API call.56        # This will raise an exception if permissions/credentials are wrong.57        # client.list_foundation_models(byOutputModality='TEXT') 58        return client, True, None # Success: client, True, no error message59    except Exception as e:60        return None, False, str(e) # Failure: None, False, error message61 62# Get the client and check its status63bedrock_client, bedrock_success, bedrock_error_msg = get_bedrock_client()64 65if not bedrock_success:66    st.error(f"Error connecting to AWS Bedrock. Please check your AWS credentials (AWS_ACCESS_KEY_ID, AWS_SECRET_ACCESS_KEY) and region (AWS_REGION) in your .env file or environment variables. Error: {bedrock_error_msg}")67    st.stop() # Stop execution if Bedrock client cannot be initialized68else:69    st.success(f"Successfully connected to AWS Bedrock in {AWS_REGION}!")70 71 72# --- LangChain Components ---73@st.cache_resource74def get_embeddings_model(_client): # Prepend underscore to tell Streamlit not to hash75    """Returns the BedrockEmbeddings model."""76    return BedrockEmbeddings(client=_client, model_id=EMBEDDING_MODEL_ID)77 78@st.cache_resource79def get_llm_model(_client): # Prepend underscore to tell Streamlit not to hash80    """Returns the Bedrock LLM model for Claude 3 Sonnet."""81    return ChatBedrock(82        client=_client,83        model_id=LLM_MODEL_ID,84        streaming=False, # <--- CHANGED: Set streaming to False85        temperature=0.1, # Lower temperature for factual accuracy in research86        model_kwargs={"max_tokens": 4000} # Claude 3 can handle larger outputs87    )88 89# --- PDF Processing and Vector Store Creation ---90def create_vector_store(pdf_file_path):91    """92    Loads PDF, chunks it contextually for mathematical papers,93    creates embeddings, and stores them in ChromaDB.94    """95    with st.spinner("Loading PDF and creating vector store..."):96        # 1. Load PDF97        loader = PyPDFLoader(pdf_file_path)98        pages = loader.load_and_split()99        st.info(f"Loaded {len(pages)} pages from the PDF.")100 101        # 2. Contextual Chunking for Mathematical Papers102        text_splitter = RecursiveCharacterTextSplitter(103            chunk_size=1500,  # Increased chunk size for math papers104            chunk_overlap=150, # Generous overlap to maintain context105            separators=[106                "\n\n",  # Prefer splitting by paragraphs107                "\n",    # Then by newlines (might break equations but less likely than fixed char)108                " ",     # Then by spaces109                "",      # Fallback110            ],111            length_function=len,112            is_separator_regex=False,113        )114        chunks = text_splitter.split_documents(pages)115        st.info(f"Split PDF into {len(chunks)} chunks.")116 117        # 3. Create Embeddings and ChromaDB118        # Pass the bedrock_client to the cached embedding model function119        embeddings = get_embeddings_model(bedrock_client) 120        vector_store = Chroma.from_documents(121            documents=chunks,122            embedding=embeddings,123            persist_directory="./chroma_db" # Persist for faster reloads (optional)124        )125        st.success("Vector store created and ready!")126        return vector_store127 128# --- RAG Chain Construction ---129def get_rag_chain(vector_store):130    """Constructs the RAG chain using LCEL."""131    retriever = vector_store.as_retriever(search_kwargs={"k": 5}) # Retrieve top 5 relevant chunks132    # Pass the bedrock_client to the cached LLM model function133    llm = get_llm_model(bedrock_client) 134 135    # Prompt Template optimized for mathematical research papers136    prompt_template = ChatPromptTemplate.from_messages(137        [138            ("system", 139             "You are an expert AI assistant specialized in analyzing and explaining mathematical research papers. "140             "Your goal is to provide precise, accurate, and concise answers based *only* on the provided context from the research paper. "141             "When answering, focus on definitions, theorems, proofs, key mathematical concepts, and experimental results. "142             "If the user asks about a mathematical notation, try to explain its meaning from the context. "143             "If the answer is not found in the context, explicitly state that you cannot find the information within the provided document. "144             "Do not invent information or make assumptions outside the given text.\n\n"145             "Context:\n{context}"),146            ("user", "{question}"),147        ]148    )149 150    rag_chain = (151        {"context": retriever, "question": RunnablePassthrough()}152        | prompt_template153        | llm154        | StrOutputParser()155    )156    return rag_chain157 158# --- Streamlit UI Main Logic ---159 160# Initialize chat history161if "messages" not in st.session_state:162    st.session_state.messages = []163 164# Initialize vector store and RAG chain165if "vector_store" not in st.session_state:166    st.session_state.vector_store = None167if "rag_chain" not in st.session_state:168    st.session_state.rag_chain = None169if "pdf_uploaded" not in st.session_state:170    st.session_state.pdf_uploaded = False171 172 173# Sidebar for PDF Upload174with st.sidebar:175    st.header("Upload PDF")176    uploaded_file = st.file_uploader(177        "Choose a PDF file",178        type="pdf",179        accept_multiple_files=False,180        key="pdf_uploader"181    )182 183    if uploaded_file and not st.session_state.pdf_uploaded:184        # Save the uploaded file temporarily185        with open("temp_doc.pdf", "wb") as f:186            f.write(uploaded_file.getbuffer())187        188        st.session_state.vector_store = create_vector_store("temp_doc.pdf")189        st.session_state.rag_chain = get_rag_chain(st.session_state.vector_store)190        st.session_state.pdf_uploaded = True191        st.success("PDF processed successfully! You can now ask questions.")192        # Clean up temporary file193        os.remove("temp_doc.pdf")194    elif st.session_state.pdf_uploaded:195        st.info("PDF already processed. Ready for questions!")196 197 198# Display chat messages from history on app rerun199for message in st.session_state.messages:200    with st.chat_message(message["role"]):201        st.markdown(message["content"])202 203# Accept user input204if prompt := st.chat_input("Ask a question about the paper..."):205    if not st.session_state.pdf_uploaded:206        st.warning("Please upload a PDF first to start asking questions.")207    else:208        # Add user message to chat history209        st.session_state.messages.append({"role": "user", "content": prompt})210        with st.chat_message("user"):211            st.markdown(prompt)212 213        # Get response from RAG chain214        with st.chat_message("assistant"):215            with st.spinner("Thinking..."):216                try:217                    # <--- CHANGED: Use invoke() instead of stream()218                    full_response = st.session_state.rag_chain.invoke(prompt) 219                    st.markdown(full_response, unsafe_allow_html=True) 220 221                    # Add assistant response to chat history222                    st.session_state.messages.append({"role": "assistant", "content": full_response})223                except Exception as e:224                    st.error(f"An error occurred during response generation: {e}")225                    st.warning("Please try again or check your AWS Bedrock access permissions.")226 227# Optional: Clear chat and uploaded PDF228if st.session_state.pdf_uploaded:229    if st.sidebar.button("Clear Chat and Upload New PDF"):230        st.session_state.messages = []231        st.session_state.vector_store = None232        st.session_state.rag_chain = None233        st.session_state.pdf_uploaded = False234        st.cache_resource.clear() # Clear streamlit caches for a clean slate235        st.rerun()