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 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 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") 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',('10', '20', '30', '50', 'All'), index=0, label_visibility = "collapsed") seed_row = st.columns([0.15, 0.45, 0.4]) with seed_row[0]: st.write("**Seed Value**") with seed_row[1]: seed_value = st.text_input("**Seed Value**", value = 20, label_visibility = "collapsed") 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") random.seed(seed_value) try: contract_list = sorted(random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count))) except: 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'])] 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') # sample = dict(zip(sample.Filename, sample.Answer)) # actual_value_list = [sample.get(contract.rsplit('.',1)[0]+'.txt', ' ') for contract in contract_list] # AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID') # AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY') # AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN') # 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','Snippet','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 = 14 else: k_value = 25 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] df['Snippet'] = [str(doc[0].page_content) for doc in doc_list] 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']) 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','Snippet' ,'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','Snippet' , '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))