Merged DEV into dev_umistry

This commit is contained in:
Umang Mistry
2024-04-22 17:06:20 +00:00
11 changed files with 7180 additions and 380 deletions
+1
View File
@@ -59,6 +59,7 @@ streamlit/contract_fields.csv
streamlit/sample.csv
streamlit/temp1.csv
streamlit/temp2.csv
streamlit/results.csv
# env
streamlit/venv
File diff suppressed because one or more lines are too long
File diff suppressed because it is too large Load Diff
+174
View File
@@ -0,0 +1,174 @@
import json
import streamlit as st
from streamlit_extras.add_vertical_space import add_vertical_space
import os
import streamlit as st
import pandas as pd
from io import StringIO
from datetime import datetime
import boto3
import util
import requests
from sf_conn import get_secret, save_to_sf
from io import StringIO, BytesIO
from sf_conn import get_client_names, insert_upload_logs
create_batch_url = 'https://lfksus2t62.execute-api.us-east-2.amazonaws.com/dev/create-batch'
REDIRECT_URI = 'https://doczy.aarete.com:8500'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@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")
_,c1= st.columns([5,1])
try:
util.setup_page(REDIRECT_URI)
except Exception as e:
st.write(f"SSO Failed = {e}")
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
try:
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
user_mail = st.session_state.user_info['mail']
except KeyError as e:
st.write("Session Expired.")
st.stop()
if user_mail in user_list:
s3_client = boto3.client('s3',
region_name="us-east-2",
)
client_list = ['doczy-ai-client-1', 'Delaware First Health, Inc.', 'Community Health Choice, Inc','CareSource Network Partners LLC',
'HealthNet of Cali', 'Oklahoma Complete Health, Inc', 'HealthFirst', 'Molina Healthcare of TX', 'AvMed', 'Arizona Care1st',
'WellCare New Jersey']
# This is the list of client fetched from Snowflake
# TODO: Need to update the streamlit code to use the client names from this list
# And use the s3 paths to save the objects for the respective client
client_list, s3_paths = get_client_names()
client_s3_paths = dict(zip(client_list, s3_paths))
client_row = st.columns([0.1, 0.8])
with client_row[0]:
st.write("**Client Name**")
with client_row[1]:
client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed")
client_bucket = client_s3_paths.get(client)
# client = 'doczy-ai-client-1'
file_row = st.columns([0.1, 0.8])
with file_row[0]:
st.write("**Upload Files**")
with file_row[1]:
file_list = st.file_uploader("Upload", type=None, accept_multiple_files=True, label_visibility = "collapsed")
add_vertical_space(2)
df = pd.DataFrame(columns=['Contract Name'])
df['Contract Name'] = file_list
file_names = []
buttons = st.columns([0.4, 0.4, 0.2])
with buttons[1]:
if st.button("Create Batch"):
myobj = { "client-bucket-name": client_bucket }
response = requests.post(create_batch_url, json = myobj)
if response.status_code >= 200 and response.status_code < 300:
try:
batch_id = json.loads(json.loads(response.text)['body'])['batch_id']
landing_zone = json.loads(json.loads(response.text)['body'])['landing_zone']
except:
st.write(myobj)
st.write(response.text)
batch_id = 'failed_cases'
landing_zone = 'contracts_landing_zone'
else:
st.write("Failed")
for uploaded_file in file_list:
stringio = BytesIO(uploaded_file.getvalue())
stringio.seek(0)
s3_client.put_object(Bucket=client_bucket, Body=stringio.getvalue(), Key=
batch_id+'/'+landing_zone+'/'+str(uploaded_file.name))
# TODO: Test this insert function with snowflake
upload_log = insert_upload_logs(batch_id, client, str(uploaded_file.name), datetime.now().strftime("%Y-%m-%d %H:%M:%S"), user_mail)
st.write(upload_log)
file_names.append(str(uploaded_file.name))
st.write(f"{batch_id} created")
st.write(f"Files uploaded to s3://{client_bucket}/{batch_id}/{landing_zone}")
# @st.cache_data
# def convert_df(df):
# return df.to_csv(index=False).encode('utf-8')
# csv = convert_df(df)
# df = df.reset_index() # make sure indexes pair with number of rows
# contract_list = []
# for index, row in df.iterrows():
# group_list = []
# if row['Unique Key']:
# group_list.append('Unique Key')
# if row['Pricing Before Carveouts']:
# group_list.append('Pricing Before Carveouts')
# if row['Contract Related']:
# group_list.append('Contract Related')
# if row['Provider']:
# group_list.append('Provider')
# if row['Timeline']:
# group_list.append('Timeline')
# if row['Carveout Indicator']:
# group_list.append('Carveout Indicator')
# if row['Carveout Methodology']:
# group_list.append('Carveout Methodology')
# entry_dict = {
# "contract_name": row['Contract Name'],
# "groups": group_list,
# "contract_source_path": "batches/batch_1/"+client+"/"+row['Contract Name']
# }
# contract_list.append(entry_dict)
# contract_list = list(df['Contract Name'])
# myobj = {
# "s3_bucket": 'doczy-dev-infra-textract',
# "batch_id": "1",
# "client_name": client,
# "username": user_mail,
# "contract_list": contract_list
# }
# buttons = st.columns([0.8, 0.2])
# with buttons[0]:
# st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
# with buttons[1]:
# if st.button("Upload to DB"):
# response = requests.post(doczy_pipeline, json = myobj)
# if response.status_code >= 200 and response.status_code < 300:
# st.write("Success")
# else:
# st.write("Failed")
# # st.write(response.text)
else:
st.write("Access Denied")
+110 -27
View File
@@ -7,10 +7,15 @@ from io import StringIO
from datetime import datetime
import boto3
import util
import requests
from sf_conn import get_client_names, get_secret, save_to_sf
REDIRECT_URI = 'http://172.29.20.126:8501'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com']
doczy_pipeline = 'https://8ir4vi1ri4.execute-api.us-east-2.amazonaws.com/dev/'
REDIRECT_URI = 'https://doczy.aarete.com:8501'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@aarete.com']
st.set_page_config(layout = "wide")
# # Sidebar contents
# with st.sidebar:
@@ -25,38 +30,55 @@ st.set_page_config(layout = "wide")
# add_vertical_space(15)
# # st.write("Doczy")
_,c1= st.columns([5,1])
try:
util.setup_page(REDIRECT_URI)
except:
st.write("SSO Failed")
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
try:
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
user_mail = st.session_state.user_info['mail']
except KeyError as e:
st.write("Session Expired.")
st. stop()
util.setup_page(REDIRECT_URI)
if st.session_state.user_info['mail'] in user_list:
if user_mail in user_list:
s3_client = boto3.client('s3',
region_name="us-east-1",
region_name="us-east-2",
)
objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract'
, Prefix="batches/batch_1/", Delimiter='/')
folder_list = []
for prefix in objects['CommonPrefixes']:
folder_list.append(prefix['Prefix'][:-1].split('/')[-1])
# # to be replaced with snowflake data
# client_list = ['doczy-ai-client-1', 'Delaware First Health, Inc.', 'Community Health Choice, Inc','CareSource Network Partners LLC',
# 'HealthNet of Cali', 'Oklahoma Complete Health, Inc', 'HealthFirst', 'Molina Healthcare of TX', 'AvMed', 'Arizona Care1st',
# 'WellCare New Jersey']
client_list, s3_paths = get_client_names()
client_s3_paths = dict(zip(client_list, s3_paths))
client_row = st.columns([0.1, 0.8])
with client_row[0]:
st.write("**Client Name**")
with client_row[1]:
client = st.selectbox('Client Name',(folder_list), index=7, label_visibility = "collapsed")
client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed")
folder_objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract'
, Prefix="batches/batch_1/"+client+"/", Delimiter='/')
client_bucket = client_s3_paths.get(client)
# # to be deleted when buckets for different clients are ready; below line is added only for testing the corresponding DAG
client_bucket = 'doczy-ai-client-1'
folder_list_2 = []
for prefix in folder_objects['CommonPrefixes']:
folder_list_2.append(prefix['Prefix'][:-1].split('/')[-1])
batch_objects = s3_client.list_objects_v2(Bucket=client_bucket
, Prefix="", Delimiter='/')
batch_list = []
for prefix in batch_objects['CommonPrefixes']:
batch_list.append(prefix['Prefix'][:-1].split('/')[-1])
path_row = st.columns([0.1, 0.8])
with path_row[0]:
st.write("**Path to folder**")
st.write("**Batch ID**")
with path_row[1]:
Directory = st.selectbox('**Path to folder**', folder_list_2, label_visibility = "collapsed")
batch_id = st.selectbox('**Batch ID**', batch_list, label_visibility = "collapsed")
checks = st.columns([0.1, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12])
with checks[0]:
@@ -78,11 +100,11 @@ if st.session_state.user_info['mail'] in user_list:
add_vertical_space(1)
df = pd.DataFrame(columns=['Request ID','Contract ID','Contract Name','Unique Key','Pricing Before Carveouts'
df = pd.DataFrame(columns=['Contract Name', 'Unique Key','Pricing Before Carveouts'
, 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology'])
file_list = []
file_objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract'
, Prefix="batches/batch_1/"+client+"/"+Directory+"/", Delimiter='/')
file_objects = s3_client.list_objects_v2(Bucket=client_bucket
, Prefix=batch_id+"/contracts_landing_zone/", Delimiter='/')
if st.button("Read the contracts from Path"):
for obj in file_objects.get('Contents',[]):
@@ -90,8 +112,8 @@ if st.session_state.user_info['mail'] in user_list:
file_list.append(obj['Key'].split('/')[-1])
df['Contract Name'] = file_list
df['Request ID'] = range(len(file_list))
df['Contract ID'] = file_list
# df['Request ID'] = range(len(file_list))
# df['Contract ID'] = file_list
df['Unique Key'] = a
df['Pricing Before Carveouts'] = b
df['Contract Related'] = c
@@ -104,23 +126,84 @@ if st.session_state.user_info['mail'] in user_list:
df.to_csv('temp1.csv', index=False)
add_vertical_space(1)
# df_copy = df.set_index(df.columns[0]).copy()
df2 = pd.read_csv('temp1.csv')
edited_df = st.data_editor(df2)
edited_df['REQUEST_USER'] = user_mail
edited_df['LATEST_FLAG BOOLEAN'] = True
edited_df['PIPELINE_KICKOFF_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
edited_df['REQUEST_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
@st.cache_data
def convert_df(df):
return df.to_csv(index=False).encode('utf-8')
csv = convert_df(edited_df)
# edited_df = edited_df.reset_index() # make sure indexes pair with number of rows
# additional_info = pd.DataFrame(columns=['REQUEST_ID','T_DRIVE_PATH','CLIENT_NAME'
# , 'GROUP_NAME', 'REQUEST_USERNAME', 'REQUEST_DATETIME'])
additional_info = pd.DataFrame(columns=['CLIENT_NAME', 'BATCH_ID', 'REQUEST_USERNAME', 'REQUEST_DATETIME'])
additional_info.loc[0] = [client, batch_id, st.session_state.user_info['mail'], datetime.now().strftime("%Y-%m-%d %H:%M:%S")]
st.write(additional_info)
contract_list = []
for index, row in edited_df.iterrows():
group_list = []
if row['Unique Key']:
group_list.append('Unique Key')
if row['Pricing Before Carveouts']:
group_list.append('Pricing Before Carveouts')
if row['Contract Related']:
group_list.append('Contract Related')
if row['Provider']:
group_list.append('Provider')
if row['Timeline']:
group_list.append('Timeline')
if row['Carveout Indicator']:
group_list.append('Carveout Indicator')
if row['Carveout Methodology']:
group_list.append('Carveout Methodology')
entry_dict = {
"contract_name": row['Contract Name'],
"groups": group_list,
"contract_source_path": batch_id+"/contracts_landing_zone/"+row['Contract Name']
}
contract_list.append(entry_dict)
myobj = {
"s3_bucket": client_bucket,
"batch_id": batch_id,
"client_name": client,
"username": user_mail,
"contract_list": contract_list
}
buttons = st.columns([0.8, 0.2])
with buttons[0]:
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
with buttons[1]:
st.button("Run Doczy.AI Pipeline")
if st.button("Run Doczy.AI Pipeline"):
# csv_buf = StringIO()
# additional_info.to_csv(csv_buf, header=True, index=False)
# csv_buf.seek(0)
# s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/request_submission.csv')
# csv_buf = StringIO()
# edited_df.to_csv(csv_buf, header=True, index=False)
# csv_buf.seek(0)
# s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv')
# try:
# save_to_sf('load_request_and_contract_submissions', request_submission_file_name = "request_submission.csv", contract_config_file_name = "contract_config.csv")
# except Exception as e:
# st.write(e)
response = requests.post(doczy_pipeline, json = myobj)
if response.status_code >= 200 and response.status_code < 300:
st.write("Success")
# st.write(myobj)
else:
st.write("Failed")
# st.write(response.text)
else:
st.write("Access Denied")
+359 -114
View File
@@ -12,12 +12,17 @@ import streamlit as st
from streamlit_extras.add_vertical_space import add_vertical_space
import os
import pandas as pd
import numpy as np
import util
import anthropic
from pydantic import BaseModel
from typing import List
import re
REDIRECT_URI = 'http://172.29.20.126:8502'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
REDIRECT_URI = 'https://doczy.aarete.com:8502'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com']
, 'vnair@aarete.com', 'kminhas@aarete.com','dculotta@aarete.com','cbull@aarete.com','sclark@aarete.com', 'fmohiuddin@aarete.com', 'hupreti@aarete.com']
st.set_page_config(layout = "wide")
# Sidebar contents
@@ -33,27 +38,82 @@ with st.sidebar:
add_vertical_space(15)
# st.write("Doczy")
util.setup_page(REDIRECT_URI)
if st.session_state.user_info['mail'] in user_list:
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
fields = fields[fields['PRIORITY'] == 'A']
fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ')
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
_,c1= st.columns([5,1])
try:
util.setup_page(REDIRECT_URI)
except:
st.write("SSO Failed")
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
try:
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
user_mail = st.session_state.user_info['mail']
except KeyError as e:
st.write("Session Expired.")
st.stop()
def file_selector(folder_path=SOURCE_DIRECTORY):
filenames = os.listdir(folder_path)
selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed")
# return os.path.join(folder_path, selected_filename)
return selected_filename
try:
sf_secrets = json.loads(get_secret())
conn = snowflake.connector.connect(
user=sf_secrets.get('user'),
password=sf_secrets.get('password'),
account="aarete-doczyai",
role = "DEVADMIN",
warehouse="DEV_XS",
database="DOCZY_DEV",
schema="STG"
)
cur = conn.cursor()
query = 'select * from "TRAINING_DATA_RAW"'
cur.execute(query)
field_values = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
# st.write(field_values)
field_values['Document_Name'] = field_values['DOCUMENT_NAME']
# field_values['Contract ID'] = field_values['CONTRACT_TITLE']
# error('table values are incorrect')
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())
fields.rename(columns={'FIELD_DESC': 'Field Name'}, inplace = True)
fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True)
fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True)
fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True)
fields.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()]
s3_client = boto3.client('s3',
region_name="us-east-2"
)
bucket = 'doczy-dev-infra-textract'
objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/")
file_list = []
for obj in objects['Contents']:
if not obj['Key'].endswith('/'):
file_list.append(obj['Key'])
contract_list = sorted(file_list)
# to be deleted later
contract_list = [contract for contract in contract_list if contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') in list(field_values['Document_Name'])]
if user_mail in user_list:
file_row = st.columns([0.2, 0.7, 0.1])
with file_row[0]:
st.write("**Contract Name**")
with file_row[1]:
# file_name = st.text_input("**Contract Name**", label_visibility = "collapsed")
file_name = file_selector()
file_name = st.selectbox('Select a file', contract_list + ['All'], label_visibility = "collapsed")
# lob_row = st.columns([0.2, 0.7, 0.1])
# with lob_row[0]:
@@ -61,24 +121,62 @@ if st.session_state.user_info['mail'] in user_list:
# with lob_row[1]:
# lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed")
llm_row = st.columns([0.2, 0.7, 0.1])
with llm_row[0]:
st.write("**Langauge Model**")
with llm_row[1]:
llm_selected = st.selectbox('Langauge Model',('Llama 2 Chat 13B', 'Llama 2 Chat 70B', 'Titan Text Express'), label_visibility = "collapsed")
field_row = st.columns([0.2, 0.7, 0.1])
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")
page_list = []
with open(os.path.join(SOURCE_DIRECTORY, file_name), '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'{file_name[:-4]}_page{page}.txt'
dict_with_pages = { 'source': { '$eq': file_path }}
page_list.append(dict_with_pages)
# 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')
# llm_row = st.columns([0.2, 0.7, 0.1])
# 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")
llm_selected = 'Claude Instant'
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?']))
# Setup bedrock
bedrock_runtime = boto3.client(
@@ -86,100 +184,242 @@ if st.session_state.user_info['mail'] in user_list:
region_name="us-east-1"
)
embeddings = BedrockEmbeddings(
client=bedrock_runtime,
model_id="amazon.titan-embed-text-v1",
)
DB = Chroma(
persist_directory=PERSIST_DIRECTORY,
embedding_function=embeddings,
client_settings=CHROMA_SETTINGS,
)
RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4})
# question_list = list(field_prompt_mapping.items())
# class OutputSchema(BaseModel):
# question_list[0]: str
# question_list[1]: str
# question_list[2]: str
# question_list[3]: str
# question_list[4]: str
# question_list[5]: str
# question_list[6]: str
# question_list[7]: str
# question_list[8]: str
# question_list[9]: str
# question_list[10]: str
# question_list[11]: str
# question_list[12]: str
# question_list[13]: str
# question_list[14]: str
# question_list[15]: str
# question_list[16]: str
# question_list[17]: str
# question_list[18]: str
# question_list[19]: str
if llm_selected == 'Titan Text Express':
LLM = Bedrock(
model_id="amazon.titan-text-express-v1",
client=bedrock_runtime,
model_kwargs={
"maxTokenCount": 4096,
"stopSequences": [],
"temperature": 0,
"topP": 1,
# class OutputList(BaseModel):
# answer: List[OutputSchema]
# question = '\\n'.join(question_list)
question = json.dumps(field_prompt_mapping)
# question_with_schema = f'{question}{OutputList.schema_json()}'
question_with_schema = question
def run_llm(bucket, file_name, llm_selected, field_values):
# with open(os.path.join(SOURCE_DIRECTORY, file_name), 'r') as infile:
# context = infile.read()
# # page_count = context.count('Start of Page No. = ')
data = s3_client.get_object(Bucket=bucket, Key=file_name)
contents = data['Body'].read()
context = contents.decode("utf-8")
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. You must answer in correct JSON format.
#
{context}
#
Question: {question}
Answer: Answer in JSON format: {{
"""
parameters = {
"maxTokenCount":1024,
"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[:6000]
prompt_data = f"""Answer the question based only on the information provided between ## and give step by step guide. You must answer in correct JSON format.
##
{context}
##
Question: {question}
Answer: Answer in JSON format: {{
"""
payload={
"prompt":"[INST]"+ prompt_data +"[/INST]",
"max_gen_len":1024,
"temperature":0.0,
"top_p":0.9
}
)
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,
}
)
else:
LLM = Bedrock(
model_id="meta.llama2-13b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
}
)
body=json.dumps(payload)
model_id="meta.llama2-70b-chat-v1"
template = """
elif llm_selected in ['Claude Instant', 'Claude 2']:
if llm_selected == 'Claude Instant':
context = context[:150000]
Use the following pieces of context to answer the question 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.
prompt_data = f"""
{context}
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. You must answer in JSON format.
Question: {question}
Answer:"""
prompt = PromptTemplate(input_variables=["context", "question"], template=template)
{context}
QA = RetrievalQA.from_chain_type(
llm=LLM,
chain_type="stuff",
retriever=RETRIEVER,
return_source_documents=True,
chain_type_kwargs={"prompt": prompt},
)
Question: {question_with_schema}
# query = "In which state or states is the Contract applicable? Answer in one or two words. State name: "
# response = QA({"query":query})
# st.write(query)
# st.write(response['result'])
# st.write("-----------")
# st.write(response)
Assistant: Answer in JSON format: {{
"""
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"
# clicked = st.button("Show Results")
df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level',
'Field Extracted Value','Imputed Value'])
field_list = list(field_prompt_mapping.keys())
query_list = [field_prompt_mapping[x] for x in field_list]
score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list]
confidence_list = []
for score in score_list:
confidence_list.append(max(d[1] for d in score))
# st.write(confidence_list)
if st.button("Show Results"):
response_list = [QA({"query":query}) for query in query_list]
answer_list = [response['result'] for response in response_list]
doc_list = [response['source_documents'] for response in response_list]
snippet_list = [str(doc[0].page_content) for doc in doc_list]
page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
# def claude_prompt_format(prompt: str) -> str:
# # Add headers to start and end of prompt
# return "\n\nHuman: " + prompt + "\n\nAssistant:"
# # Call Claude model
# def call_claude(prompt):
# prompt_config = {
# "prompt": claude_prompt_format(prompt),
# "max_tokens_to_sample": 4096,
# "temperature": 0.5,
# "top_k": 250,
# "top_p": 0.5,
# "stop_sequences": [],
# }
# body = json.dumps(prompt_config)
# modelId = "anthropic.claude-instant-v1"
# accept = "application/json"
# contentType = "application/json"
# response = bedrock_runtime.invoke_model(
# body=body, modelId=modelId, accept=accept, contentType=contentType
# )
# response_body = json.loads(response.get("body").read())
# results = response_body.get("completion")
# return results
# prompt = SECOND_PROMPT
# result = call_claude(prompt)
# st.write(result)
df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number',
'Field Extracted Value','Imputed Value'])
# field_list = list(field_prompt_mapping.keys())
# query_list = [field_prompt_mapping[x] for x in field_list]
# st.write(field_prompt_mapping)
# st.write(prompt_data)
try:
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']
except:
response_text = "failed"
raw_response_text = response_text
response_text = response_text.strip()
try:
if response_text.split("{",1)[1].strip()[0] == '"':
response_text = "{" + response_text.split("{",1)[1]
else:
response_text = "{" + response_text
except:
response_text = "{" + response_text
if len(response_text.split("}",1)) > 1:
if response_text.rsplit("}",1)[0].strip()[-1] == '"':
response_text = response_text.rsplit("}",1)[0] + "}"
else:
response_text = response_text.rstrip(",")
response_text = response_text + "}"
try:
response_dict = json.loads(response_text)
except:
response_dict = {"Test value": "Failed to extract"}
# st.write(response_dict)
field_list = list(response_dict.keys())
answer_list = list(response_dict.values())
# st.write(answer_list)
location_list = [context.find(answer) if isinstance(answer, str) and answer != "" else -1 for answer in answer_list]
snippet_list = [' '.join(context[:location].split('.')[-4:]) + ' ' + ' '.join(context[location:].split('. ')[:5]
) if location != -1 else ' ' for location in location_list]
page_no_list = [" " 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] for location in location_list]
# st.write(location_list)
page_no_list = [re.search(r'\d+', page).group() if page != " " and re.search(r'\d+', page) is not None else "" for page in page_no_list]
# response_list = [QA({"query":query}) for query in query_list]
# answer_list = [response['result'] for response in response_list]
# doc_list = [response['source_documents'] for response in response_list]
# snippet_list = [str(doc[0].page_content) for doc in doc_list]
# page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
df['Field Name'] = field_list
df['Contract Name'] = file_name
df['Snippet'] = snippet_list
df['Page Number'] = page_no_list
df['Confidence Level'] = confidence_list
# df['Confidence Level'] = ' '
df['Field Extracted Value'] = answer_list
df.to_csv('temp2.csv', index=False)
df = pd.merge(df, fields[['Field Name', 'SF_DB_COL_NAME']], how ='left', on ='Field Name')
document_name = [x for x in list(field_values['Document_Name']) if not pd.isna(x) and file_name.rsplit('/',1)[1].replace(' MU','').replace(
'_MU','').replace('.txt','') in x][0]
field_values = field_values[field_values['Document_Name'] == document_name].head(1).transpose().reset_index()
field_values.columns = ['SF_DB_COL_NAME', 'Actual Value']
df = pd.merge(df, field_values, how ='left', on ='SF_DB_COL_NAME')
df = df[['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number', 'Field Extracted Value', 'Actual Value','Imputed Value']]
return df, raw_response_text
raw_response_text = '{"Test value": "Failed to extract"}'
if st.button("Show Results"):
if file_name == 'All':
df_1 = pd.DataFrame(columns=['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number'
, 'Field Extracted Value', 'Actual Value','Imputed Value'])
for contract in contract_list:
print(contract)
df, raw_response_text = run_llm(bucket, contract, llm_selected, field_values)
df_1 = pd.concat([df_1, df], ignore_index = True)
else:
df_1, raw_response_text = run_llm(bucket, file_name, llm_selected, field_values)
df_1.to_csv('temp2.csv', index=False)
df2 = pd.read_csv('temp2.csv')
df2['Imputed Value'] = ''
@@ -193,12 +433,17 @@ if st.session_state.user_info['mail'] in user_list:
buttons = st.columns(3)
with buttons[0]:
st.button("Save All Imputations")
with buttons[1]:
# st.button("Save All Imputations")
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
with buttons[1]:
# st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
st.write("")
with buttons[2]:
st.button("Kickoff Database Integration")
add_vertical_space(20)
st.write(raw_response_text)
else:
st.write("Access Denied")
+232
View File
@@ -0,0 +1,232 @@
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 os
import pandas as pd
import util
REDIRECT_URI = 'https://doczy.aarete.com:8502'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@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")
util.setup_page(REDIRECT_URI)
_,c1= st.columns([5,1])
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
if st.session_state.user_info['mail'] in user_list:
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
fields = fields[fields['PRIORITY'] == 'A']
fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ')
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
def file_selector(folder_path=SOURCE_DIRECTORY):
filenames = os.listdir(folder_path)
selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed")
# return os.path.join(folder_path, selected_filename)
return selected_filename
file_row = st.columns([0.2, 0.7, 0.1])
with file_row[0]:
st.write("**Contract Name**")
with file_row[1]:
# file_name = st.text_input("**Contract Name**", label_visibility = "collapsed")
file_name = file_selector()
# lob_row = st.columns([0.2, 0.7, 0.1])
# with lob_row[0]:
# st.write("**LOB**")
# with lob_row[1]:
# lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed")
llm_row = st.columns([0.2, 0.7, 0.1])
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")
page_list = []
with open(os.path.join(SOURCE_DIRECTORY, file_name), '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'{file_name[:-4]}_page{page}.txt'
dict_with_pages = { 'source': { '$eq': file_path }}
page_list.append(dict_with_pages)
# 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"
)
embeddings = BedrockEmbeddings(
client=bedrock_runtime,
model_id="amazon.titan-embed-text-v1",
)
DB = Chroma(
persist_directory=PERSIST_DIRECTORY,
embedding_function=embeddings,
client_settings=CHROMA_SETTINGS,
)
RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4})
# 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
template = """
Use the following pieces of context to answer the question 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}
Answer:"""
prompt = PromptTemplate(input_variables=["context", "question"], template=template)
QA = RetrievalQA.from_chain_type(
llm=LLM,
chain_type="stuff",
retriever=RETRIEVER,
return_source_documents=True,
chain_type_kwargs={"prompt": prompt},
)
# query = "In which state or states is the Contract applicable? Answer in one or two words. State name: "
# response = QA({"query":query})
# st.write(query)
# st.write(response['result'])
# st.write("-----------")
# st.write(response)
# clicked = st.button("Show Results")
df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level',
'Field Extracted Value','Imputed Value'])
field_list = list(field_prompt_mapping.keys())
query_list = [field_prompt_mapping[x] for x in field_list]
score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list]
confidence_list = []
for score in score_list:
confidence_list.append(max(d[1] for d in score))
# st.write(confidence_list)
if st.button("Show Results"):
response_list = [QA({"query":query}) for query in query_list]
answer_list = [response['result'] for response in response_list]
doc_list = [response['source_documents'] for response in response_list]
snippet_list = [str(doc[0].page_content) for doc in doc_list]
page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
df['Field Name'] = field_list
df['Contract Name'] = file_name
df['Snippet'] = snippet_list
df['Page Number'] = page_no_list
df['Confidence Level'] = confidence_list
df['Field Extracted Value'] = answer_list
df.to_csv('temp2.csv', index=False)
df2 = pd.read_csv('temp2.csv')
df2['Imputed Value'] = ''
edited_df = st.data_editor(df2)
@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")
else:
st.write("Access Denied")
+586 -236
View File
@@ -11,16 +11,23 @@ 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, save_to_sf
from io import StringIO
REDIRECT_URI = 'http://172.29.20.126:8503'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
REDIRECT_URI = 'https://doczy.aarete.com:8503'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com']
, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@aarete.com']
st.set_page_config(layout = "wide")
# Sidebar contents
@@ -36,46 +43,203 @@ with st.sidebar:
add_vertical_space(15)
# st.write("Doczy")
util.setup_page(REDIRECT_URI)
_,c1= st.columns([5,1])
# util.setup_page(REDIRECT_URI)
try:
util.setup_page(REDIRECT_URI)
except:
st.write("SSO Failed")
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
try:
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
user_mail = st.session_state.user_info['mail']
except KeyError as e:
st.write("Session Expired.")
st. stop()
try:
sf_secrets = json.loads(get_secret())
conn = snowflake.connector.connect(
user=sf_secrets.get('user'),
password=sf_secrets.get('password'),
account="aarete-doczyai",
role = "DEVADMIN",
warehouse="DEV_XS",
database="DOCZY_DEV",
schema="STG"
)
cur = conn.cursor()
query = 'select * from "TRAINING_DATA_RAW"'
cur.execute(query)
field_values = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
# st.write(field_values)
field_values['Document_Name'] = field_values['DOCUMENT_NAME']
# field_values['Contract ID'] = field_values['CONTRACT_TITLE']
# error('table values are incorrect')
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)
field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:')]
st.write("Local copy of TRAINING_DATA_RAW table loaded")
try:
# error('table is not updated')
query = 'select * from "BUSINESS_CONFIG"'
cur.execute(query)
fields = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
fields.rename(columns={'FIELD_NAME': 'Field Name'}, inplace = True)
fields.rename(columns={'QUESTION': 'Interrogation Question?'}, inplace = True)
fields.rename(columns={'SF_COL_NAME': 'SF_DB_COL_NAME'}, inplace = True)
fields = fields[~fields['SF_DB_COL_NAME'].str.endswith('_PG', na=None)]
fields['Field Name'] = fields['SF_DB_COL_NAME']
# fields.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['Field Name'] = fields['SF_DB_COL_NAME']
fields = fields[~fields['Field Name'].isnull()]
st.write("Local copy of BUSINESS_CONFIG table loaded")
try:
query = 'select * from "TRAINING_ATTEMPT_LOGS"'
cur.execute(query)
history = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
history.rename(columns={'FIELD_NAME': 'Field Name'}, inplace = True)
history.rename(columns={'CONTRACTS_TESTED': '# Contracts Tested'}, inplace = True)
history.rename(columns={'USERNAME': 'Username'}, inplace = True)
history.rename(columns={'DATE_TIME': 'Date/Time'}, inplace = True)
history.rename(columns={'ACCURACY': 'Accuracy'}, inplace = True)
history.rename(columns={'ATTEMPT_NUM': 'Attempt #'}, inplace = True)
history = history[['Field Name', '# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #']]
except:
history = pd.read_csv('history.csv')
st.write("Local copy of TRAINING_ATTEMPT_LOGS table loaded")
if st.session_state.user_info['mail'] 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_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 and 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 and Contract Related':
fields = fields[fields['PRIORITY'].isin(['A','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?']))
mode_row = st.columns([0.15, 0.45, 0.4])
with mode_row[0]:
st.write("**Mode**")
with mode_row[1]:
mode = st.selectbox('Mode',('Single field - Non Empty values', 'Multiple fields'), index=0, label_visibility = "collapsed")
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")
if mode != 'Multiple fields':
field = st.selectbox('Field Name',sorted(set(field_prompt_mapping.keys())), index=0, label_visibility = "collapsed")
else:
field = st.multiselect('Field Name',sorted(set(field_prompt_mapping.keys())), sorted(
set(field_prompt_mapping.keys())), label_visibility = "collapsed")
field_prompt_mapping = {key: field_prompt_mapping[key] for key in field}
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")
contract_count = st.selectbox('Contract count',('1', '10', '20', '30', '50', '100', '200','Part-1', 'Part-2', 'All'), index=1, label_visibility = "collapsed")
s3_client = boto3.client('s3',
region_name="us-east-2"
)
bucket = 'doczy-dev-infra-textract'
objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/")
file_list = []
for obj in objects['Contents']:
if not obj['Key'].endswith('/'):
file_list.append(obj['Key'])
# print(os.path.join('RAW_DOCUMENTS', obj['Key'].rsplit('/',1)[1]))
# s3_client.download_file('doczy-dev-infra-textract', obj['Key'], os.path.join('RAW_DOCUMENTS', obj['Key'].rsplit('/',1)[1]))
contract_list = sorted(file_list)
# contract_list = sorted(os.listdir(SOURCE_DIRECTORY))
if mode == 'Single field - Non Empty values':
# column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0]
# field_values = field_values[~field_values[column_name].isnull()]
field_values = field_values[~field_values[field].isnull()]
# df = pd.DataFrame({'col':contract_list})
# st.write(df)
# st.write(len(contract_list))
# contract_list = [contract for contract in contract_list if contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') in list(field_values['Document_Name'])]
contract_list = [str(contract)[:-4]+'.pdf' for contract in contract_list]
contract_list = [contract for contract in contract_list if contract.rsplit('/',1)[1] in list(field_values['Document_Name'])]
# st.write(len(contract_list))
if contract_count == 'All':
contract_count = len(contract_list)
if contract_count == 'Part-1':
contract_list = np.array_split(contract_list, 2)[0]
contract_count = len(contract_list)
if contract_count == 'Part-2':
contract_list = np.array_split(contract_list, 2)[1]
contract_count = len(contract_list)
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']:
if contract_count in ['10', '20', '30', '50', '100', '200']:
st.write("**Seed Value**")
elif contract_count == '1':
st.write("**Contract Name**")
with seed_row[1]:
if contract_count in ['10', '20', '30', '50']:
if contract_count in ['10', '20', '30', '50', '100', '200']:
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)))
# contract_list = sorted(random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count)))
contract_list = sorted(random.choices(contract_list, k=int(contract_count)))
elif contract_count == '1':
contract_name = st.selectbox('Contract Name', (contract_list), label_visibility = "collapsed")
contract_list = [contract_name]
@@ -85,42 +249,36 @@ if st.session_state.user_info['mail'] in user_list:
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")
llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude 3 - Haiku', 'Claude 3 - Sonnet', 'Claude Instant'
, 'Llama 2 Chat 70B', 'Titan Text Express'), index=3, label_visibility = "collapsed")
st.write("**Prompt**")
sequence_input = field_prompt_mapping.get(field)
if mode == 'Multiple fields':
sequence_input = json.dumps(field_prompt_mapping)
else:
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")
# st.button("Save Prompt")
with prompt_row[0]:
prompt = st.text_area("**Prompt**", sequence_input, height = 150, label_visibility = "collapsed")
prompt = st.text_area("**Prompt**", sequence_input, height = 100, 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)
# column_name = fields.loc[fields['Field Name'] == field, 'SF_DB_COL_NAME'].iloc[0]
# column_list = ['Document_Name', column_name]
# # 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.rename(columns={'Document_Name': 'Contract Name'}, inplace=True)
field_values = field_values.drop_duplicates(subset='Contract Name', keep="first").sort_values('Contract Name')
field_values = field_values[field_values['Contract Name'].isin([contract.rsplit('/',1)[1] for contract in contract_list])]
# Setup bedrock
bedrock_runtime = boto3.client(
@@ -128,224 +286,408 @@ if st.session_state.user_info['mail'] in user_list:
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
# question = prompt
if mode == 'Multiple fields':
prompt_dict = json.loads(prompt)
prompt_dict_pg = prompt_dict | {str(k)+'_PG': "On which page can I find answer to the question - "+str(
v) for k, v in prompt_dict.items()}
question = json.dumps(dict(sorted(prompt_dict_pg.items())))
else:
k_value = 20
question = json.dumps({field: prompt, str(field)+'_PG': "On which page can I find answer to the question - "+str(prompt)})
# st.write(question)
question_with_schema = question
attempt = 0
def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, history):
if st.button("Test Configuration"):
# 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' ,'New Extracted value','Confidence Level','Snippet','New Page Number'
, 'Revised Prompt', 'Result'])
field_list = []
answer_list = []
doc_list = []
response_list = []
score_list = []
snippet_list = []
page_no_list = []
contract_list_f = []
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)
for contract in contract_list:
# with open(os.path.join(SOURCE_DIRECTORY, contract[:-4]+'.txt'), 'r') as infile:
# context = infile.read()
data = s3_client.get_object(Bucket=bucket, Key=str(contract)[:-4]+'.txt')
contents = data['Body'].read()
context = contents.decode("utf-8")
# st.write(question_with_schema)
# 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. You must answer in correct JSON format.
#
{context}
#
Question: {question}
Answer: Answer in JSON format: {{"""
parameters = {
"maxTokenCount":2048,
"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[:6000]
prompt_data = f"""Answer the question based only on the information provided between ## and give step by step guide. You must answer in correct JSON format.
##
{context}
##
Question: {question}
Answer: Answer in JSON format: {{"""
payload={
"prompt":"[INST]"+ prompt_data +"[/INST]",
"max_gen_len":2048,
"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', 'Claude 3 - Haiku', 'Claude 3 - Sonnet']:
if llm_selected == 'Claude Instant':
context = context[:175000]
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. You must answer in correct JSON format.
{context}
Question: {question_with_schema}
Assistant: Answer in JSON format: {{"""
if llm_selected == "Claude 2":
model_id = "anthropic.claude-v2:1"
body = json.dumps(
{"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT,
"max_tokens_to_sample": 2048,
"temperature":0.0,
"top_p":1,
"top_k":250,
"stop_sequences":[anthropic.HUMAN_PROMPT]
})
elif llm_selected in ["Claude 3 - Haiku", "Claude 3 - Sonnet"]:
if llm_selected == "Claude 3 - Haiku":
model_id = 'anthropic.claude-3-haiku-20240307-v1:0'
else:
model_id = 'anthropic.claude-3-sonnet-20240229-v1:0'
body = json.dumps({
"anthropic_version": "bedrock-2023-05-31",
"max_tokens": 2048,
"messages": [
{
"role": "user",
"content": [
{
"type": "text",
"text":anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT
}
]
}
],
"temperature": 0.0
}
)
else:
model_id = "anthropic.claude-instant-v1"
body = json.dumps(
{"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT,
"max_tokens_to_sample": 2048,
"temperature":0.0,
"top_p":1,
"top_k":250,
"stop_sequences":[anthropic.HUMAN_PROMPT]
})
try:
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']
elif llm_selected in ['Claude 3 - Haiku', 'Claude 3 - Sonnet']:
response_text = response_body['content'][0]['text']
except:
response_text = "failed"
raw_response_text = response_text
response_text = response_text.strip()
try:
if response_text.split("{",1)[1].strip()[0] == '"':
response_text = "{" + response_text.split("{",1)[1]
else:
response_text = "{" + response_text
except:
response_text = "{" + response_text
if len(response_text.split("}",1)) > 1:
if response_text.rsplit("}",1)[0].strip()[-1] == '"':
response_text = response_text.rsplit("}",1)[0] + "}"
else:
response_text = response_text.rstrip(",")
response_text = response_text + "}"
try:
response_dict = json.loads(response_text)
except:
if mode == 'Multiple fields':
response_dict = {"Test field": "Failed to extract"}
else:
response_dict = {field:response_text.strip("{").strip("}")}
# if mode == 'Multiple fields':
# field_l = list(response_dict.keys())
# answer_l = list(response_dict.values())
# else:
# field_l = [field]
# # answer = response_dict.get(field, " ")
# try:
# if isinstance(response_dict, dict):
# answer_l = list(response_dict.values())[:1]
# else:
# answer_l = list(response_dict)[:1]
# except:
# answer_l = [response_dict]
if mode == 'Multiple fields':
field_l = list(response_dict.keys())
answer_l = list(response_dict.values())
field_dict = {k: v for k, v in response_dict.items() if not k.endswith('_PG')}
page_dict = {k: v for k, v in response_dict.items() if k.endswith('_PG')}
page_dict = {k[:-3]: v for k, v in response_dict.items()}
field_l = list(field_dict.keys())
answer_l = list(field_dict.values())
page_no_l = [page_dict.get(x, "") for x in field_l]
else:
field_l = [field]
answer_l = [list(response_dict.values())[0]]
try:
page_no_l = [list(response_dict.values())[1]]
except:
page_no_l = ['']
field_list.extend(field_l)
answer_list.extend(answer_l)
contract_list_f.extend([contract]*len(field_l))
# 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 ""
try:
location_l = [context.find(a, context.find("Start of Page No. = "+str(p))) if isinstance(
a, str) and a != "" else -1 for a, p in zip(answer_l, page_no_l)]
except:
location_l = [context.find(answer) if isinstance(answer, str) and answer != "" else -1 for answer in answer_l]
snippet_l = [' '.join(context[:location].split('.')[-4:]) + ' ' + ' '.join(context[location:].split('. ')[:5]
) if location != -1 else ' ' for location in location_l]
# page_no_l = [" " 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] for location in location_l]
# # st.write(location_list)
# page_no_l = [re.search(r'\d+', page).group() if page != " " and re.search(r'\d+', page) is not None else "" for page in page_no_l]
snippet_list.extend(snippet_l)
page_no_list.extend(page_no_l)
df['Field Name'] = field_list
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]
# 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
try:
answer_list = [str(answer).strip("\n").strip().strip("[").strip("]").strip("{").strip("}").strip('"').rstrip('"').strip(
' ') if answer is not None else None for answer in answer_list]
if 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 3 - Haiku', 'Claude 3 - Sonnet', 'Claude Instant']:
answer_list = [answer if "do not have" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "do not see" 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 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 if "N/A" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "don't know" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "do not see" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Not specified" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Don't know" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "don't see" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "don't have" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Does not apply" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "Nothing found" 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]
except Exception as e:
st.write(e)
df['Contract Name'] = contract_list
# to be deleted later
df['Contract Name'] = [contract.replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list]
# contract_list_f = [contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list_f]
contract_list_f = [contract.rsplit('/',1)[1] for contract in contract_list_f]
df['Contract Name'] = contract_list_f
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['Confidence Level'] = ' '
df['Snippet'] = snippet_list
df['New Page Number'] = page_no_list
df['New Page Number'] = df['New Page Number'].apply(lambda x: re.search(r'\d+', x).group(
) if isinstance(x, str) and re.search(r'\d+', x) is not None else " ")
df['Revised Prompt'] = [prompt] * len(contract_list_f)
df = pd.merge(df, field_values, how ='left', on ='Contract Name')
df = pd.merge(df, fields[['Field Name', 'SF_DB_COL_NAME']], how ='left', on ='Field Name')
field_values_2 = pd.DataFrame(columns=['Contract Name', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Original Page Number'])
for file_name in contract_list:
# document_name = [x for x in list(field_values['Contract Name']) if not pd.isna(x) and file_name.rsplit('/',1)[1].replace(' MU','').replace(
# '_MU','').replace('.txt','') in x][0]
document_name = [x for x in list(field_values['Contract Name']) if not pd.isna(x) and file_name.rsplit('/',1)[1] in x][0]
field_values_1 = field_values[field_values['Contract Name'] == document_name].head(1).transpose().reset_index()
field_values_1.columns = ['SF_DB_COL_NAME', 'Actual Value Stored']
field_values_p1 = field_values_1[~field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')]
field_values_p2 = field_values_1[field_values_1['SF_DB_COL_NAME'].str.endswith('_PG')]
field_values_p2.columns = ['SF_DB_COL_NAME', 'Original Page Number']
field_values_p2["SF_DB_COL_NAME"] = field_values_p2["SF_DB_COL_NAME"].str.replace("_PG", "")
field_values_1 = pd.merge(field_values_p1, field_values_p2, how ='left', on =['SF_DB_COL_NAME'])
field_values_1['Contract Name'] = document_name
field_values_2 = pd.concat([field_values_2, field_values_1], ignore_index = True)
# st.write(df)
# st.write(field_values_2)
df = pd.merge(df, field_values_2, how ='right', on =['Contract Name', 'SF_DB_COL_NAME'])
if mode == 'Multiple fields':
df = df[df['SF_DB_COL_NAME'].isin(list(field_prompt_mapping.keys()))]
else:
df = df[df['SF_DB_COL_NAME'].isin([field])]
df = df.drop_duplicates(subset=['SF_DB_COL_NAME', 'Contract Name'], keep="first")
df['Original Page Number'] = df['Original Page Number'].apply(lambda x: re.search(r'\d+', x).group(
) if isinstance(x, str) and re.search(r'\d+', x) is not None else " ")
df['Raw value 2'] = df['New Extracted value']
df_date = df[df['SF_DB_COL_NAME'].str.contains('_DT', na=False)]
df_others = df[~df['SF_DB_COL_NAME'].str.contains('_DT', na=False)]
df_date['Actual Value Stored'] = pd.to_datetime(df_date['Actual Value Stored'],errors='coerce').dt.strftime('%Y-%m-%d').fillna(" ")
df_date['New Extracted value'] = pd.to_datetime(df_date['New Extracted value'],errors='coerce').dt.strftime('%Y-%m-%d').fillna(" ")
df = pd.concat([df_date, df_others], ignore_index = True)
df.sort_values(['SF_DB_COL_NAME', 'Contract Name'], inplace=True)
df.fillna(" ", inplace=True)
df['Raw value 3'] = df['New Extracted value']
df['Actual Value Stored'] = df['Actual Value Stored'].apply(lambda x: x.strip() if isinstance(x, str) else '')
df['New Extracted value'] = df['New Extracted value'].apply(lambda x: x.strip() if isinstance(x, str) else '')
actual_value_list = list(df['Actual Value Stored'])
actual_value_list = [answer if str(answer) != "12 months" else "1 year" for answer in actual_value_list]
actual_value_list = [answer if str(answer) != "Fifth" else "5" for answer in actual_value_list]
actual_value_list = [answer if str(answer) != "Seventh" else "7" for answer in actual_value_list]
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)]
answer_list = [answer if str(answer) != "one-year" else "1 year" for answer in answer_list]
answer_list = [answer if str(answer) != "one year" else "1 year" for answer in answer_list]
answer_list = [answer if str(answer) != "one" else "1 year" for answer in answer_list]
answer_list = [answer if str(answer) != "one (1) year" else "1 year" for answer in answer_list]
answer_list = [answer if str(answer) != "twelve" else "1 year" for answer in answer_list]
answer_list = [answer if str(answer) != "XI" else "11" for answer in answer_list]
answer_list = [answer if str(answer) != "Third" else "3" for answer in answer_list]
answer_list = [answer if str(answer) != "Six" else "6" for answer in answer_list]
actual_value_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in actual_value_list]
answer_list = [s.replace('-', '').replace(' ', '').replace('[', '').replace(']', '').lower() for s in answer_list]
result_list = [(i in j) or (j in i) if isinstance(i, str) and isinstance(
j, str) and ((i != '') == (j != '')) else False 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']]
# df = df[~df['Contract ID'].isnull()]
# df['Contract ID'] = contract_list_f
df['Contract ID'] = range(len(actual_value_list))
df = df[['Contract Name','Contract ID', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Raw value', 'Raw value 2', 'Raw value 3'
, 'New Extracted value','Confidence Level','Snippet','Original Page Number', '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]
if mode == 'Multiple fields':
field = field_group
history.loc[len(history.index)] = [field, str(contract_count), st.session_state.user_info['mail'], 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()
return df, history, attempt, raw_response_text
raw_response_text = ''
df = pd.DataFrame(columns=['Contract Name','Contract ID', 'Actual Value Stored','New Extracted value','Confidence Level'
,'Snippet','New Page Number', 'Revised Prompt', 'Result'])
if st.button("Test Configuration"):
df, history, attempt, raw_response_text = run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, history)
df_1 = df[['Contract Name','Contract ID', 'Actual Value Stored','New Extracted value','Confidence Level'
,'Snippet','Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
df.to_csv('results.csv', index=False)
history.to_csv('history.csv', index=False)
# s3_client.upload_file('results.csv', bucket, key)
csv_buf = StringIO()
df_1.to_csv(csv_buf, header=True, index=False)
csv_buf.seek(0)
s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/results.csv')
csv_buf = StringIO()
history.tail(1).to_csv(csv_buf, header=True, index=False)
csv_buf.seek(0)
s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/history.csv')
df = df[['Contract Name','Contract ID', 'SF_DB_COL_NAME', 'Actual Value Stored','New Extracted value','Confidence Level'
,'Snippet','Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
try:
df = pd.read_csv('results.csv')
df['Result'] = df['Result'].astype('str')
except:
df = pd.DataFrame(columns=['Contract Name' ,'New Extracted value','Confidence Level','Snippet','New Page Number'
, 'Revised Prompt', 'Result'])
history = pd.read_csv('history.csv')
st.dataframe(df)
st.dataframe(history)
# @st.cache_data
# def convert_df(df):
# return df.to_csv(index=False).encode('utf-8')
@@ -360,9 +702,17 @@ if st.session_state.user_info['mail'] in user_list:
# with buttons[2]:
# st.button("Kickoff Database Integration")
st.write(column_name)
add_vertical_space(20)
st.write(field)
if mode != 'Multiple fields':
st.write(fields.loc[fields['SF_DB_COL_NAME'] == field, 'Field Name'].iloc[0])
st.write(len(contract_list))
st.write(raw_response_text)
try:
save_to_sf('load_training_results', training_results_file_name = "results.csv", attempt_logs_file_name = "history.csv")
except:
st.write("running locally")
else:
st.write("Access Denied")
+371
View File
@@ -0,0 +1,371 @@
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
import util
REDIRECT_URI = 'http://172.29.20.126:8503'
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@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")
# 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")
+37 -3
View File
@@ -1,16 +1,50 @@
import streamlit as st
import msal
import requests
import boto3
from botocore.exceptions import ClientError
import json
# Replace with your own values
CLIENT_ID = 'effafe90-7ed7-43a3-ab03-19a0be2f1758'
CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD'
# CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD'
# TENANT_ID = ''
AUTHORITY = 'https://login.microsoftonline.com/organizations/'
SCOPE = ['User.Read']
# REDIRECT_URI = 'http://localhost:8501'
REDIRECT_URI = 'https://171.29.20.126:8501'
# Initialize boto3 client to interact with AWS Secrets Manager
def get_secret():
secret_name = "doczy-sso-azure-app-key"
region_name = "us-east-2"
# Create a Secrets Manager client
# session = boto3.session.Session()
# client = session.client(
# service_name='secretsmanager',
# region_name=region_name
# )
client = boto3.client('secretsmanager', region_name=region_name)
try:
get_secret_value_response = client.get_secret_value(
SecretId=secret_name
)
except ClientError as e:
raise e
secret = get_secret_value_response['SecretString']
secret = json.loads(secret)['CLIENT_SECRET']
return secret
CLIENT_SECRET = get_secret()
app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET)
@@ -19,7 +53,7 @@ def get_auth_url(REDIRECT_URI):
auth_url = app.get_authorization_request_url(SCOPE, redirect_uri=REDIRECT_URI)
return auth_url
@st.cache_data
def get_token_from_code(auth_code, REDIRECT_URI):
app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET)
result = app.acquire_token_by_authorization_code(auth_code, scopes=SCOPE, redirect_uri=REDIRECT_URI)
+140
View File
@@ -0,0 +1,140 @@
# Use this code snippet in your app.
# If you need more information about configurations
# or implementing the sample code, visit the AWS docs:
# https://aws.amazon.com/developer/language/python/
import boto3
from botocore.exceptions import ClientError
import http.client
import base64
import ast
import snowflake.connector
import json
def get_secret():
secret_name = "doczy_dev_db_creds"
region_name = "us-east-2"
# Create a Secrets Manager client
session = boto3.session.Session()
client = session.client(
service_name='secretsmanager',
region_name=region_name
)
try:
get_secret_value_response = client.get_secret_value(
SecretId=secret_name
)
except ClientError as e:
# For a list of exceptions thrown, see
# https://docs.aws.amazon.com/secretsmanager/latest/apireference/API_GetSecretValue.html
raise e
secret = get_secret_value_response['SecretString']
return secret
# get_secret()
# TODO: This function needs to be changed to accept Kwargs
# The function name should be more generic and cofnigurable
def save_to_sf(dag_name, **kwargs):
mwaa_env_name = 'doczy-dev-infra-mwaa'
dag_name = dag_name
mwaa_cli_command = 'dags trigger'
# Create the client with the specified profile
session = boto3.Session()
client = session.client('mwaa', region_name='us-east-2')
# get web token
mwaa_cli_token = client.create_cli_token(
Name=mwaa_env_name
)
conn = http.client.HTTPSConnection(mwaa_cli_token['WebServerHostname'])
# This section passes the payload to the MWAA CLI
# The file parameters should be added dynamically in streamlit, once the file names are passed while triggering the dag, the data will be ingested
# training_results_file = "training_results_sample.csv"
# attempt_logs_file = "attempt_logs_sample.csv"
# conf = "{\"" + "training_results_file_name" + "\":\"" + {training_results_file} + "\", \"" + "attempt_logs_file_name" + "\":\"" + {attempt_logs_file} + "\"}".format(training_results_file=training_results_file, attempt_logs_file=attempt_logs_file)
conf = json.dumps(kwargs)
payload = mwaa_cli_command + " " + dag_name + " --conf '{}'".format(conf)
headers = {
'Authorization': 'Bearer ' + mwaa_cli_token['CliToken'],
'Content-Type': 'text/plain'
}
conn.request("POST", "/aws_mwaa/cli", payload, headers)
res = conn.getresponse()
data = res.read()
dict_str = data.decode("UTF-8")
mydata = ast.literal_eval(dict_str)
return payload
# save_to_sf("2024-03-13T17-36_prompt_results.csv", "2024-03-14T23-19_history.csv")
# Create snowflake connection with the secrets fetched from secrets manager for the schema passed as an argument
def get_snowflake_conn(schema):
secret = get_secret()
secret_dict = eval(secret)
conn = snowflake.connector.connect(
user=secret_dict['user'],
password=secret_dict['password'],
account=secret_dict['account_alias'],
warehouse=secret_dict['warehouse'],
database=secret_dict['database'],
schema=schema
)
return conn
def get_client_names():
"""
Input: None
Output: client_names, s3_paths lists
"""
try:
# Get conn from snowflake_conn function for STG schema
conn = get_snowflake_conn('STG')
cursor = conn.cursor()
# Query to get client names and their s3_paths
query = "SELECT DISTINCT client_name, bucket_name FROM STG.CLIENT_LOGS"
cursor.execute(query)
# Create 2 lists from the query results
client_names = []
s3_paths = []
for row in cursor:
client_names.append(row[0])
s3_paths.append(row[1]) # Extracting the bucket name from the s3 path
cursor.close()
conn.close()
return client_names, s3_paths
except Exception as e:
print(f"Error while fetching client names from snowflake: {e}")
return None, None
def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload_user):
"""
Input: batch_id, client_name, file_name, upload_datetime, upload_user
Output: status of the insert query
"""
try:
conn = get_snowflake_conn('STG')
cursor = conn.cursor()
query = f"INSERT INTO STG.CONTRACT_UPLOAD_LOGS (BATCH_ID, CLIENT_NAME, FILE_NAME, UPLOAD_DATETIME, UPLOAD_USER) VALUES ('{batch_id}', '{client_name}', '{file_name}', '{upload_datetime}', '{upload_user}')"
cursor.execute(query)
cursor.close()
conn.close()
return 'Log inserted successfully'
except Exception as e:
return e