432 lines
18 KiB
Python
432 lines
18 KiB
Python
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
|
|
import numpy as np
|
|
from datetime import datetime
|
|
import random
|
|
import os
|
|
import dateutil
|
|
import util
|
|
import anthropic
|
|
import re
|
|
import snowflake.connector
|
|
from sf_conn import get_secret
|
|
|
|
|
|
REDIRECT_URI = 'https://doczy.aarete.com:8503'
|
|
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
|
|
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
|
|
, 'vnair@aarete.com']
|
|
|
|
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")
|
|
|
|
try:
|
|
util.setup_page(REDIRECT_URI)
|
|
_,c1= st.columns([4,1])
|
|
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
|
|
user_mail = st.session_state.user_info['mail']
|
|
except:
|
|
user_mail = 'maamseek@aarete.com'
|
|
|
|
try:
|
|
sf_secrets = get_secret()
|
|
conn = snowflake.connector.connect(
|
|
user=sf_secrets['user'],
|
|
password=sf_secrets['password'],
|
|
account="aarete-doczyai",
|
|
role = sf_secrets['ROLE'],
|
|
warehouse=sf_secrets['warehouse'],
|
|
database=sf_secrets['database'],
|
|
schema="STG"
|
|
)
|
|
cur = conn.cursor()
|
|
|
|
query = 'select * from "TRAINING_DATA_RAW"'
|
|
cur.execute(query)
|
|
field_values = pd.DataFrame(cur.fetchall())
|
|
field_values['Contract ID'] = field_values['Document_Name']
|
|
|
|
except:
|
|
field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True)
|
|
field_values.rename(columns={'(internal) Document Name': 'Document_Name'}, inplace = True)
|
|
field_values.rename(columns={'(Internal) Carveout ID': 'Contract ID'}, inplace = True)
|
|
st.write("conn failed")
|
|
|
|
try:
|
|
query = 'select * from "PROMPT_CONFIG"'
|
|
cur.execute(query)
|
|
fields = pd.DataFrame(cur.fetchall())
|
|
field_values.rename(columns={'FIELD_DESC': 'Field Name'}, inplace = True)
|
|
field_values.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True)
|
|
field_values.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True)
|
|
field_values.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True)
|
|
field_values.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True)
|
|
error('table is empty')
|
|
except:
|
|
fields = pd.read_csv('contract_fields.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()]
|
|
|
|
|
|
if user_mail in user_list:
|
|
|
|
field_row = st.columns([0.15, 0.45, 0.4])
|
|
with field_row[0]:
|
|
st.write("**Field Group**")
|
|
with field_row[1]:
|
|
field_group = st.selectbox('Field Group',('Unique Key', 'Contract Related', 'Pricing Before Carveouts - I'
|
|
, 'Pricing Before Carveouts - II', 'Carveout Indicator, Code Type and Code #s - I'
|
|
, 'Carveout Indicator, Code Type and Code #s - II', 'Carveout Indicator, Code Type and Code #s - III'
|
|
, 'Optimize Carving Indic.', 'Carveout Method - I', 'Carveout Method - II', 'Provider'
|
|
, 'Timeline'), label_visibility = "collapsed")
|
|
|
|
|
|
# priorty column will be relaced by group_id in snowflake db
|
|
if field_group == 'Unique Key':
|
|
fields = fields[fields['PRIORITY'] == 'A']
|
|
elif field_group == 'Contract Related':
|
|
fields = fields[fields['PRIORITY'] == 'C']
|
|
elif field_group == 'Pricing Before Carveouts - I':
|
|
fields = fields[fields['PRIORITY'] == 'B']
|
|
fields = np.array_split(fields, 2)[0]
|
|
elif field_group == 'Pricing Before Carveouts - II':
|
|
fields = fields[fields['PRIORITY'] == 'B']
|
|
fields = np.array_split(fields, 2)[1]
|
|
elif field_group == 'Carveout Indicator, Code Type and Code #s - I':
|
|
fields = fields[fields['PRIORITY'] == 'F']
|
|
fields = np.array_split(fields, 3)[0]
|
|
elif field_group == 'Carveout Indicator, Code Type and Code #s - II':
|
|
fields = fields[fields['PRIORITY'] == 'F']
|
|
fields = np.array_split(fields, 3)[1]
|
|
elif field_group == 'Carveout Indicator, Code Type and Code #s - III':
|
|
fields = fields[fields['PRIORITY'] == 'F']
|
|
fields = np.array_split(fields, 3)[2]
|
|
elif field_group == 'Carveout Methodology - I':
|
|
fields = fields[fields['PRIORITY'] == 'G']
|
|
fields = np.array_split(fields, 4)[0]
|
|
elif field_group == 'Carveout Methodology - II':
|
|
fields = fields[fields['PRIORITY'] == 'G']
|
|
fields = np.array_split(fields, 4)[1]
|
|
elif field_group == 'Carveout Method - III':
|
|
fields = fields[fields['PRIORITY'] == 'G']
|
|
fields = np.array_split(fields, 4)[2]
|
|
elif field_group == 'Carveout Method - IV':
|
|
fields = fields[fields['PRIORITY'] == 'G']
|
|
fields = np.array_split(fields, 4)[3]
|
|
elif field_group == 'Provider':
|
|
fields = fields[fields['PRIORITY'] == 'D']
|
|
elif field_group == 'Timeline':
|
|
fields = fields[fields['PRIORITY'] == 'E']
|
|
|
|
fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ')
|
|
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['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 70B'
|
|
, 'Titan Text Express'), index=1, 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")
|
|
|
|
|
|
column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0]
|
|
column_list = ['Document_Name', 'Contract 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={'Document_Name': 'Contract Name', column_name: 'Actual Value Stored'
|
|
, 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",
|
|
)
|
|
|
|
# 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 #'])
|
|
|
|
|
|
question = prompt
|
|
question_with_schema = question
|
|
attempt = 0
|
|
|
|
if st.button("Test Configuration"):
|
|
|
|
answer_list = []
|
|
snippet_list = []
|
|
page_no_list = []
|
|
attempt = attempt + 1
|
|
|
|
for contract in contract_list:
|
|
with open(os.path.join(SOURCE_DIRECTORY, contract[:-4]+'.txt'), 'r') as infile:
|
|
context = infile.read()
|
|
|
|
# Add "You must answer in correct JSON format."
|
|
# Add Answer in JSON format: {{
|
|
if llm_selected == "Titan Text Express":
|
|
context = context[:16000]
|
|
prompt_data = f"""Answer the question based only on the information provided between ## and give step by step guide.
|
|
#
|
|
{context}
|
|
#
|
|
|
|
Question: {question}
|
|
Answer:"""
|
|
parameters = {
|
|
"maxTokenCount":512,
|
|
"stopSequences":[],
|
|
"temperature":0,
|
|
"topP":0.9
|
|
}
|
|
|
|
body = json.dumps({"inputText": prompt_data, "textGenerationConfig": parameters})
|
|
model_id = "amazon.titan-text-express-v1" # change this to use a different version from the model provider
|
|
|
|
elif llm_selected == 'Llama 2 Chat 70B':
|
|
context = context[:7000]
|
|
prompt_data = f"""Answer the question based only on the information provided between ## and give step by step guide.
|
|
##
|
|
{context}
|
|
##
|
|
|
|
Question: {question}
|
|
Answer:"""
|
|
payload={
|
|
"prompt":"[INST]"+ prompt_data +"[/INST]",
|
|
"max_gen_len":512,
|
|
"temperature":0.0,
|
|
"top_p":0.9
|
|
}
|
|
body=json.dumps(payload)
|
|
model_id="meta.llama2-70b-chat-v1"
|
|
|
|
elif llm_selected in ['Claude Instant', 'Claude 2']:
|
|
|
|
prompt_data = f"""
|
|
|
|
Human: Use the following pieces of context to provide a concise answer to the questions at the end. If you don't know the answer, just say that you don't know, don't try to make up an answer.
|
|
|
|
{context}
|
|
|
|
Question: {question_with_schema}
|
|
|
|
Assistant:"""
|
|
body = json.dumps(
|
|
{"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT,
|
|
"max_tokens_to_sample": 1024,
|
|
"temperature":0.0,
|
|
"top_p":1,
|
|
"top_k":250,
|
|
"stop_sequences":[anthropic.HUMAN_PROMPT]
|
|
})
|
|
if llm_selected == "Claude 2":
|
|
model_id = "anthropic.claude-v2:1"
|
|
else:
|
|
model_id = "anthropic.claude-instant-v1"
|
|
|
|
response = bedrock_runtime.invoke_model(
|
|
body=body,
|
|
modelId=model_id,
|
|
accept="application/json",
|
|
contentType="application/json"
|
|
)
|
|
response_body = json.loads(response.get("body").read())
|
|
|
|
if llm_selected == "Titan Text Express":
|
|
response_text = response_body.get("results")[0].get("outputText")
|
|
elif llm_selected == 'Llama 2 Chat 70B':
|
|
response_text = response_body['generation']
|
|
elif llm_selected in ['Claude Instant', 'Claude 2']:
|
|
response_text = response_body['completion']
|
|
|
|
# st.write(response_text)
|
|
try:
|
|
response_text = "{" + response_text.split("{",1)[1]
|
|
response_text = response_text.split("}",1)[0] + "}"
|
|
response_dict = json.loads(response_text)
|
|
except:
|
|
response_dict = {field:response_text}
|
|
|
|
answer = response_dict.get(field, " ")
|
|
answer_list.append(answer)
|
|
location = context.find(answer) if isinstance(answer, str) and answer != "" else -1
|
|
snippet = ' '.join(context[:location].split()[-25:]) + ' ' + ' '.join(context[location:].split()[:30]) if location != -1 else ' '
|
|
page_no = " " if location == -1 else context[:location].rsplit("Start of Page No. = ", 1)[1] if len(context[:location].rsplit(
|
|
"Start of Page No. = ", 1)) > 1 else context[:location].rsplit("Start of Page No. = ", 1)[0]
|
|
page_no = re.search(r'\d+', page_no).group() if page_no != " " and re.search(r'\d+', page_no) is not None else ""
|
|
snippet_list.append(snippet)
|
|
page_no_list.append(page_no)
|
|
|
|
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 "do not have" not in str(answer) else " " for answer in answer_list]
|
|
answer_list = [answer if "does not specify" not in str(answer) else " " for answer in answer_list]
|
|
answer_list = [answer if "does not explicitly" 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]
|
|
|
|
|
|
# to be deleted later
|
|
contract_list = [contract.replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list]
|
|
|
|
df['Contract Name'] = contract_list
|
|
|
|
df['New Extracted value'] = answer_list
|
|
df['Confidence Level'] = ' '
|
|
df['Snippet'] = snippet_list
|
|
df['New Page Number'] = page_no_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'])
|
|
if 'Date' in field:
|
|
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'
|
|
,'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))
|
|
|
|
else:
|
|
st.write("Access Denied")
|
|
|
|
|