Team Ai
Apppublic

andreped/ReferenceBot

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
main.py128 linesDownload Raw Back to knowledge_gpt
1import os2 3import streamlit as st4from langchain.chat_models import AzureChatOpenAI5 6from knowledge_gpt.components.sidebar import sidebar7from knowledge_gpt.core.caching import bootstrap_caching8from knowledge_gpt.core.chunking import chunk_file9from knowledge_gpt.core.embedding import embed_files10from knowledge_gpt.core.parsing import read_file11from knowledge_gpt.core.qa import query_folder12from knowledge_gpt.ui import display_file_read_error13from knowledge_gpt.ui import is_file_valid14from knowledge_gpt.ui import is_query_valid15from knowledge_gpt.ui import wrap_doc_in_html16 17st.set_page_config(page_title="ReferenceBot", page_icon="📖", layout="wide")18 19# add all secrets into environmental variables20if os.path.exists(21    os.path.dirname(os.path.abspath(__file__)) + "/../.streamlit/secrets.toml"22):  # to avoid redundant print by calling st.secrets23    for key, value in st.secrets.items():24        os.environ[key] = value25 26 27def main():28    EMBEDDING = "openai"29    VECTOR_STORE = "faiss"30    MODEL_LIST = ["gpt-3.5-turbo", "gpt-4"]31 32    # Uncomment to enable debug mode33    # MODEL_LIST.insert(0, "debug")34 35    st.header("📖ReferenceBot")36 37    # Enable caching for expensive functions38    bootstrap_caching()39 40    sidebar()41 42    uploaded_file = st.file_uploader(43        "Upload a pdf, docx, or txt file",44        type=["pdf", "docx", "txt"],45        help="Scanned documents are not supported yet!",46    )47 48    model: str = st.selectbox("Model", options=MODEL_LIST)  # type: ignore49 50    with st.expander("Advanced Options"):51        return_all_chunks = st.checkbox("Show all chunks retrieved from vector search")52        show_full_doc = st.checkbox("Show parsed contents of the document")53 54    if not uploaded_file:55        st.stop()56 57    try:58        file = read_file(uploaded_file)59    except Exception as e:60        display_file_read_error(e, file_name=uploaded_file.name)61 62    chunked_file = chunk_file(file, chunk_size=300, chunk_overlap=0)63 64    if not is_file_valid(file):65        st.stop()66 67    with st.spinner("Indexing document... This may take a while⏳"):68        folder_index = embed_files(69            files=[chunked_file],70            embedding=EMBEDDING if model != "debug" else "debug",71            vector_store=VECTOR_STORE if model != "debug" else "debug",72            deployment=os.environ["ENGINE_EMBEDDING"],73            model=os.environ["ENGINE"],74            openai_api_key=os.environ["OPENAI_API_KEY"],75            openai_api_base=os.environ["OPENAI_API_BASE"],76            openai_api_type="azure",77            chunk_size=1,78        )79 80    with st.form(key="qa_form"):81        query = st.text_area("Ask a question about the document")82        submit = st.form_submit_button("Submit")83 84    if show_full_doc:85        with st.expander("Document"):86            # Hack to get around st.markdown rendering LaTeX87            st.markdown(f"<p>{wrap_doc_in_html(file.docs)}</p>", unsafe_allow_html=True)88 89    if submit:90        if not is_query_valid(query):91            st.stop()92 93        # Output Columns94        answer_col, sources_col = st.columns(2)95 96        with st.spinner("Setting up AzureChatOpenAI bot..."):97            llm = AzureChatOpenAI(98                openai_api_base=os.environ["OPENAI_API_BASE"],99                openai_api_version=os.environ["OPENAI_API_VERSION"],100                deployment_name=os.environ["ENGINE"],101                openai_api_key=os.environ["OPENAI_API_KEY"],102                openai_api_type="azure",103                temperature=0,104            )105 106        with st.spinner("Querying folder to get result..."):107            result = query_folder(108                folder_index=folder_index,109                query=query,110                return_all=return_all_chunks,111                llm=llm,112            )113 114        with answer_col:115            st.markdown("#### Answer")116            st.markdown(result.answer)117 118        with sources_col:119            st.markdown("#### Sources")120            for source in result.sources:121                st.markdown(source.page_content)122                st.markdown(source.metadata["source"])123                st.markdown("---")124 125 126if __name__ == "__main__":127    main()128