+2
-1
@@ -55,10 +55,11 @@ streamlit/DB/
|
||||
streamlit/RAW_DOCUMENTS/
|
||||
streamlit/SOURCE_DOCUMENTS/
|
||||
streamlit/contract_field_values.csv
|
||||
streamlit/contract_fields.csv
|
||||
# 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
@@ -0,0 +1,177 @@
|
||||
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']
|
||||
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:
|
||||
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()
|
||||
|
||||
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 = 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 = '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 }
|
||||
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='doczy-dev-infra-raw-data-ingestion', Body=stringio.getvalue(), Key=
|
||||
# 'config_interface/'+str(uploaded_file.name))
|
||||
s3_client.put_object(Bucket=client, 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}/{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
@@ -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']
|
||||
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
@@ -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']
|
||||
|
||||
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/contract-text-file/")
|
||||
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")
|
||||
|
||||
|
||||
@@ -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']
|
||||
|
||||
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")
|
||||
|
||||
|
||||
|
||||
+537
-239
@@ -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']
|
||||
|
||||
st.set_page_config(layout = "wide")
|
||||
# Sidebar contents
|
||||
@@ -36,46 +43,201 @@ 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:
|
||||
_,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()
|
||||
|
||||
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
|
||||
fields = fields[fields['PRIORITY'].isin(['A','C'])]
|
||||
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()]
|
||||
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
|
||||
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:
|
||||
|
||||
field_row = st.columns([0.15, 0.45, 0.4])
|
||||
with field_row[0]:
|
||||
st.write("**Field Name**")
|
||||
st.write("**Field Group**")
|
||||
with field_row[1]:
|
||||
field = st.selectbox('Field Name',sorted(set(field_prompt_mapping.keys())), index=0, label_visibility = "collapsed")
|
||||
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]:
|
||||
if mode != 'Multiple fields':
|
||||
st.write("**Field Name**")
|
||||
with field_row[1]:
|
||||
if mode != 'Multiple fields':
|
||||
field = st.selectbox('Field Name',sorted(set(field_prompt_mapping.keys())), index=0, label_visibility = "collapsed")
|
||||
else:
|
||||
field = sorted(set(field_prompt_mapping.keys()))
|
||||
|
||||
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/contract-text-file/")
|
||||
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 = [contract for contract in contract_list if contract.rsplit('/',1)[1].replace('.txt','.pdf') 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,41 +247,34 @@ 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 Instant', 'Llama 2 Chat 70B'
|
||||
, 'Titan Text Express'), index=1, 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')
|
||||
|
||||
# Setup bedrock
|
||||
@@ -128,224 +283,359 @@ 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=contract)
|
||||
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']:
|
||||
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: {{"""
|
||||
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]
|
||||
})
|
||||
if llm_selected == "Claude 2":
|
||||
model_id = "anthropic.claude-v2:1"
|
||||
else:
|
||||
model_id = "anthropic.claude-instant-v1"
|
||||
|
||||
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:
|
||||
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]
|
||||
|
||||
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')
|
||||
# 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 = [answer.strip("\n").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 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 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.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:
|
||||
print('post processing failed')
|
||||
|
||||
answer_list = list(df['New Extracted value'])
|
||||
df['Actual Value Stored'] = pd.to_datetime(df['Actual Value Stored'],errors='coerce').dt.date
|
||||
# df['Contract ID'] = contract_list_f
|
||||
df['Contract ID'] = range(len(contract_list_f))
|
||||
|
||||
# to be deleted later
|
||||
# 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].replace('.txt','.pdf') for contract in contract_list_f]
|
||||
|
||||
df['Contract Name'] = contract_list_f
|
||||
df['New Extracted value'] = answer_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, 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].replace('.txt','.pdf') 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)
|
||||
|
||||
df = pd.merge(df, field_values_2, how ='left', on =['Contract Name', 'SF_DB_COL_NAME'])
|
||||
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_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['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'])
|
||||
result_list = [i==j for i, j in zip(actual_value_list, answer_list)]
|
||||
answer_list = list(df['New Extracted value'])
|
||||
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[['Contract Name','Contract ID', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Raw value', '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 +650,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")
|
||||
|
||||
@@ -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']
|
||||
|
||||
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")
|
||||
|
||||
|
||||
@@ -9,7 +9,7 @@ CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD'
|
||||
|
||||
AUTHORITY = 'https://login.microsoftonline.com/organizations/'
|
||||
SCOPE = ['User.Read']
|
||||
# REDIRECT_URI = 'http://localhost:8501'
|
||||
REDIRECT_URI = 'https://171.29.20.126:8501'
|
||||
|
||||
|
||||
app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET)
|
||||
@@ -19,7 +19,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)
|
||||
|
||||
@@ -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, s3_bucket_path FROM STG.CLIENT_CONFIG"
|
||||
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].replace('s3://','').split('/')[0]) # 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
|
||||
Reference in New Issue
Block a user