import json import boto3 from langchain.prompts import PromptTemplate from langchain.embeddings.bedrock import BedrockEmbeddings from langchain.llms.bedrock import Bedrock from langchain_community.vectorstores import Chroma from constants import ( CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY, USER_LIST, ) from langchain.chains import RetrievalQA import streamlit as st from streamlit_extras.add_vertical_space import add_vertical_space import pandas as pd from datetime import datetime import random import os import dateutil import util REDIRECT_URI = "http://172.29.20.126:8503" user_list = USER_LIST st.set_page_config(layout="wide") # Sidebar contents with st.sidebar: st.title("Doczy.AI ™") st.markdown( """ ## About This app extracts data from contracts """ ) add_vertical_space(15) # st.write("Doczy") # util.setup_page(REDIRECT_URI) # if st.session_state.user_info['mail'] in user_list: if "maamseek@aarete.com" in user_list: fields = pd.read_csv( "contract_fields.csv", encoding="unicode_escape", skipinitialspace=True ) fields = fields[fields["PRIORITY"].isin(["A", "C"])] field_values = pd.read_csv( "contract_field_values.csv", encoding="unicode_escape", skipinitialspace=True ) fields = fields.drop_duplicates(subset="Field Name", keep="first").sort_values( "Field Name" ) fields = fields[~fields["Field Name"].isnull()] field_prompt_mapping = dict( zip(fields["Field Name"], fields["Interrogation Question?"]) ) field_row = st.columns([0.15, 0.45, 0.4]) with field_row[0]: st.write("**Field Name**") with field_row[1]: field = st.selectbox( "Field Name", sorted(set(field_prompt_mapping.keys())), index=0, label_visibility="collapsed", ) contract_count_row = st.columns([0.15, 0.45, 0.4]) with contract_count_row[0]: st.write("**# of Contracts**") with contract_count_row[1]: contract_count = st.selectbox( "Contract count", ("1", "10", "20", "30", "50", "All"), index=1, label_visibility="collapsed", ) seed_row = st.columns([0.15, 0.45, 0.4]) contract_list = sorted(os.listdir(SOURCE_DIRECTORY)) # to be deleted later contract_list = [ contract for contract in contract_list if contract.replace(" MU", "").replace("_MU", "").replace(".txt", "") in list(field_values["(internal) Document Name"]) ] with seed_row[0]: if contract_count in ["10", "20", "30", "50"]: st.write("**Seed Value**") elif contract_count == "1": st.write("**Contract Name**") with seed_row[1]: if contract_count in ["10", "20", "30", "50"]: seed_value = st.text_input( "**Seed Value**", value=20, label_visibility="collapsed" ) random.seed(seed_value) contract_list = sorted( random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count)) ) elif contract_count == "1": contract_name = st.selectbox( "Contract Name", (contract_list), label_visibility="collapsed" ) contract_list = [contract_name] llm_row = st.columns([0.15, 0.45, 0.4]) with llm_row[0]: st.write("**Langauge Model**") with llm_row[1]: llm_selected = st.selectbox( "Langauge Model", ( "Claude 2", "Claude Instant", "Llama 2 Chat 13B", "Llama 2 Chat 70B", "Titan Text Express", ), label_visibility="collapsed", ) st.write("**Prompt**") sequence_input = field_prompt_mapping.get(field) prompt_row = st.columns([0.8, 0.2]) with prompt_row[1]: if st.button("Clear Prompt"): sequence_input = "" if st.button("Back to default"): prompt = sequence_input st.button("Save Prompt") with prompt_row[0]: prompt = st.text_area( "**Prompt**", sequence_input, height=150, label_visibility="collapsed" ) page_list_all = [] for contract in contract_list: page_list = [] with open( os.path.join(SOURCE_DIRECTORY, contract[:-4] + ".txt"), "r" ) as infile: text = infile.read() page_count = text.count("Start of Page No. = ") for page in range(page_count + 1): file_path = "SOURCE_DOCUMENTS\\" + f"{contract[:-4]}_page{page}.txt" dict_with_pages = {"source": {"$eq": file_path}} page_list.append(dict_with_pages) page_list_all.append(page_list) contract_txt_mapping = dict(zip(contract_list, page_list_all)) column_name = fields.loc[fields["Field Name"] == field, "SF_DB_COL_NAME"].iloc[0] column_list = ["(internal) Document Name", "(Internal) Carveout ID", column_name] if column_name + "_PG" in list(field_values.columns): column_list.append(column_name + "_PG") field_values = field_values[column_list] field_values.rename( columns={ "(internal) Document Name": "Contract Name", column_name: "Actual Value Stored", "(Internal) Carveout ID": "Contract ID", column_name + "_PG": "Original Page Number", }, inplace=True, ) field_values = field_values.drop_duplicates( subset="Contract Name", keep="first" ).sort_values("Contract Name") # Setup bedrock bedrock_runtime = boto3.client( service_name="bedrock-runtime", region_name="us-east-1", ) # Define the retreiver # load the vectorstore if "EMBEDDINGS" not in st.session_state: EMBEDDINGS = BedrockEmbeddings( client=bedrock_runtime, model_id="amazon.titan-embed-text-v1", ) st.session_state.EMBEDDINGS = EMBEDDINGS if "DB" not in st.session_state: DB = Chroma( persist_directory=PERSIST_DIRECTORY, embedding_function=st.session_state.EMBEDDINGS, client_settings=CHROMA_SETTINGS, ) st.session_state.DB = DB # if "RETRIEVER" not in st.session_state: # # { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} } # RETRIEVER = DB.as_retriever(search_kwargs={"filter": { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }, "k": 2}) # st.session_state.RETRIEVER = RETRIEVER # if "LLM" not in st.session_state: if llm_selected == "Titan Text Express": LLM = Bedrock( model_id="amazon.titan-text-express-v1", client=bedrock_runtime, model_kwargs={ "maxTokenCount": 512, "stopSequences": [], "temperature": 0, "topP": 1, }, ) elif llm_selected == "Llama 2 Chat 70B": LLM = Bedrock( model_id="meta.llama2-70b-chat-v1", client=bedrock_runtime, model_kwargs={ "max_gen_len": 512, "temperature": 0, # "topP": 0.9, }, ) elif llm_selected == "Llama 2 Chat 13B": LLM = Bedrock( model_id="meta.llama2-13b-chat-v1", client=bedrock_runtime, model_kwargs={ "max_gen_len": 512, "temperature": 0, # "topP": 0.9, }, ) elif llm_selected == "Claude Instant": LLM = Bedrock( model_id="anthropic.claude-instant-v1", client=bedrock_runtime, model_kwargs={ # "max_tokens_to_sample": 512, "temperature": 0, # "topP": 0.9, }, ) elif llm_selected == "Claude 2": LLM = Bedrock( model_id="anthropic.claude-v2:1", client=bedrock_runtime, model_kwargs={ # "max_tokens_to_sample": 512, "temperature": 0, # "topP": 0.9, }, ) st.session_state["LLM"] = LLM # if "QA" not in st.session_state: # prompt, memory = model_memory() # QA = RetrievalQA.from_chain_type( # llm=LLM, # chain_type="stuff", # retriever=RETRIEVER, # return_source_documents=True, # chain_type_kwargs={"prompt": prompt, "memory": memory}, # ) # st.session_state["QA"] = QA # df = pd.DataFrame(columns=['Contract Name','Contract ID','Actual Value Stored','New Extracted value','Confidence Level','Snippet', # 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']) df = pd.DataFrame( columns=[ "Contract Name", "Raw value", "New Extracted value", "Confidence Level", "Snippet1", "Snippet2", "Snippet3", "Snippet4", "Snippet5", "Snippet6", "Snippet7", "Snippet8", "Snippet9", "Snippet10", "New Page Number", "Revised Prompt", "Result", ] ) try: history = pd.read_csv("history.csv") except: history = pd.DataFrame( columns=[ "Field Name", "# Contracts Tested", "Username", "Date/Time", "Accuracy", "Attempt #", ] ) attempt = 0 if llm_selected in ["Llama 2 Chat 13B", "Llama 2 Chat 70B"]: k_value = 10 else: k_value = 20 if st.button("Test Configuration"): answer_list = [] doc_list = [] response_list = [] score_list = [] attempt = attempt + 1 for page_list in page_list_all: RETRIEVER = st.session_state.DB.as_retriever( search_kwargs={"filter": {"$or": page_list}, "k": k_value} ) QA = RetrievalQA.from_chain_type( llm=st.session_state["LLM"], chain_type="stuff", retriever=RETRIEVER, return_source_documents=True, # chain_type_kwargs={"prompt": prompt, "memory": None}, ) score = st.session_state.DB.similarity_search_with_relevance_scores( prompt, k=4, filter={"$or": page_list} ) score_list.append(max(d[1] for d in score)) response = QA(prompt) answer, docs = response["result"], response["source_documents"] answer_list.append(answer) doc_list.append(docs) response_list.append(response) df["Raw value"] = answer_list # post-processing if "Date" in field: date_list = [] for answer in answer_list: try: extracted_date = dateutil.parser.parse( str(answer).replace('"', ""), fuzzy=True ).date() except: extracted_date = " " date_list.append(extracted_date) answer_list = date_list elif llm_selected in ["Llama 2 Chat 13B", "Llama 2 Chat 70B"]: answer_list = [answer.rstrip(".") for answer in answer_list] answer_list = [ answer if "I don't know" not in str(answer) else " " for answer in answer_list ] answer_list = [ answer if "N/A" not in str(answer) else " " for answer in answer_list ] answer_list = [ answer if "does not contain" not in str(answer) else " " for answer in answer_list ] answer_list = [ answer if "None" not in str(answer) else " " for answer in answer_list ] answer_list = [ answer if "Not specified in the contract" not in str(answer) else " " for answer in answer_list ] answer_list = [ answer if "Not applicable" not in str(answer) else " " for answer in answer_list ] elif llm_selected in ["Claude 2", "Claude Instant"]: answer_list = [ ( answer if "Unfortunately, I do not have enough context" not in str(answer) else " " ) for answer in answer_list ] answer_list = [answer.rstrip(".") for answer in answer_list] else: answer_list = [answer.rstrip(".") for answer in answer_list] # answer_list = [str(x).rsplit(':',1)[0] if len(str(x).rsplit(':',1)) < 2 else str(x).rsplit(':',1)[1] for x in answer_list] df["Contract Name"] = contract_list # to be deleted later df["Contract Name"] = [ contract.replace(" MU", "").replace("_MU", "").replace(".txt", "") for contract in contract_list ] df["New Extracted value"] = answer_list df["Confidence Level"] = [round(score, 2) for score in score_list] Snippet = [] count = 0 for i in range(int(contract_count)): for j in range(10): try: content = str(doc_list[i][j].page_content) except: content = " " Snippet.append(content) # df['Snippet1'] = [str(doc[0].page_content) for doc in doc_list] df["Snippet1"] = Snippet[: int(contract_count)] df["Snippet2"] = Snippet[int(contract_count) : 2 * int(contract_count)] df["Snippet3"] = Snippet[2 * int(contract_count) : 3 * int(contract_count)] df["Snippet4"] = Snippet[3 * int(contract_count) : 4 * int(contract_count)] df["Snippet5"] = Snippet[4 * int(contract_count) : 5 * int(contract_count)] df["Snippet6"] = Snippet[5 * int(contract_count) : 6 * int(contract_count)] df["Snippet7"] = Snippet[6 * int(contract_count) : 7 * int(contract_count)] df["Snippet8"] = Snippet[7 * int(contract_count) : 8 * int(contract_count)] df["Snippet9"] = Snippet[8 * int(contract_count) : 9 * int(contract_count)] df["Snippet10"] = Snippet[9 * int(contract_count) :] df["New Page Number"] = [ int(str(doc[0].metadata["source"]).rsplit("_page")[1].replace(".txt", "")) + 1 for doc in doc_list ] df["Revised Prompt"] = [prompt] * len(contract_list) df = pd.merge(df, field_values, how="left", on="Contract Name") answer_list = list(df["New Extracted value"]) df["Actual Value Stored"] = pd.to_datetime( df["Actual Value Stored"], errors="coerce" ).dt.date df.fillna(" ", inplace=True) actual_value_list = list(df["Actual Value Stored"]) result_list = [i == j for i, j in zip(actual_value_list, answer_list)] df["Result"] = [str(x) for x in result_list] df = df[~df["Contract ID"].isnull()] if "Original Page Number" in df.columns: df = df[ [ "Contract Name", "Contract ID", "Actual Value Stored", "Raw value", "New Extracted value", "Confidence Level", "Snippet1", "Snippet2", "Snippet3", "Snippet4", "Snippet5", "Snippet6", "Snippet7", "Snippet8", "Snippet9", "Snippet10", "Original Page Number", "New Page Number", "Revised Prompt", "Result", ] ] else: df = df[ [ "Contract Name", "Contract ID", "Actual Value Stored", "Raw value", "New Extracted value", "Confidence Level", "Snippet1", "Snippet2", "Snippet3", "Snippet4", "Snippet5", "Snippet6", "Snippet7", "Snippet8", "Snippet9", "Snippet10", "New Page Number", "Revised Prompt", "Result", ] ] try: accuracy = round( sum(bool(x) for x in result_list) * 100 / len(list(df["Result"])), 2 ) except: accuracy = "NA" history.loc[len(history.index)] = [ field, str(contract_count), None, datetime.now().strftime("%Y-%m-%d %H:%M:%S"), accuracy, attempt, ] # df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False) history.to_csv("history.csv", index=False) # df_copy = df.set_index(df.columns[0]).copy() # df_2_copy = history.set_index(history.columns[0]).copy() st.dataframe(df) st.dataframe(history) # @st.cache_data # def convert_df(df): # return df.to_csv(index=False).encode('utf-8') # csv = convert_df(edited_df) # buttons = st.columns(3) # with buttons[0]: # st.button("Save All Imputations") # with buttons[1]: # st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') # with buttons[2]: # st.button("Kickoff Database Integration") st.write(column_name) st.write(len(contract_list)) else: st.write("Access Denied")