added business users for sso
This commit is contained in:
@@ -8,7 +8,7 @@ from datetime import datetime
|
||||
import boto3
|
||||
import util
|
||||
|
||||
REDIRECT_URI = 'http://localhost:8501'
|
||||
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']
|
||||
st.set_page_config(layout = "wide")
|
||||
|
||||
@@ -14,9 +14,10 @@ import os
|
||||
import pandas as pd
|
||||
import util
|
||||
|
||||
REDIRECT_URI = 'http://localhost:8502'
|
||||
REDIRECT_URI = 'http://172.29.20.126:8502'
|
||||
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com'
|
||||
, 'piragavarapu@aarete.com', 'umistry@aarete.com']
|
||||
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
|
||||
, 'vnair@aarete.com']
|
||||
|
||||
st.set_page_config(layout = "wide")
|
||||
# Sidebar contents
|
||||
|
||||
+323
-313
@@ -15,9 +15,14 @@ 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'
|
||||
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
|
||||
, 'vnair@aarete.com']
|
||||
|
||||
st.set_page_config(layout = "wide")
|
||||
|
||||
# Sidebar contents
|
||||
with st.sidebar:
|
||||
st.title("Doczy.AI ™")
|
||||
@@ -31,336 +36,341 @@ with st.sidebar:
|
||||
add_vertical_space(15)
|
||||
# st.write("Doczy")
|
||||
|
||||
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
|
||||
fields = fields[fields['PRIORITY'].isin(['A','C'])]
|
||||
field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True)
|
||||
fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
|
||||
fields = fields[~fields['Field Name'].isnull()]
|
||||
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
|
||||
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'].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")
|
||||
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")
|
||||
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])
|
||||
seed_row = st.columns([0.15, 0.45, 0.4])
|
||||
|
||||
|
||||
contract_list = sorted(os.listdir(SOURCE_DIRECTORY))
|
||||
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')
|
||||
|
||||
# sample = dict(zip(sample.Filename, sample.Answer))
|
||||
# actual_value_list = [sample.get(contract.rsplit('.',1)[0]+'.txt', ' ') for contract in contract_list]
|
||||
|
||||
# AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID')
|
||||
# AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY')
|
||||
# AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN')
|
||||
|
||||
# Setup bedrock
|
||||
bedrock_runtime = boto3.client(
|
||||
service_name="bedrock-runtime",
|
||||
region_name="us-east-1",
|
||||
)
|
||||
|
||||
# Define the retreiver
|
||||
# load the vectorstore
|
||||
if "EMBEDDINGS" not in st.session_state:
|
||||
EMBEDDINGS = BedrockEmbeddings(
|
||||
client=bedrock_runtime,
|
||||
model_id="amazon.titan-embed-text-v1",
|
||||
)
|
||||
st.session_state.EMBEDDINGS = EMBEDDINGS
|
||||
|
||||
if "DB" not in st.session_state:
|
||||
DB = Chroma(
|
||||
persist_directory=PERSIST_DIRECTORY,
|
||||
embedding_function=st.session_state.EMBEDDINGS,
|
||||
client_settings=CHROMA_SETTINGS,
|
||||
)
|
||||
st.session_state.DB = DB
|
||||
|
||||
# if "RETRIEVER" not in st.session_state:
|
||||
# # { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }
|
||||
# RETRIEVER = DB.as_retriever(search_kwargs={"filter": { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }, "k": 2})
|
||||
# st.session_state.RETRIEVER = RETRIEVER
|
||||
|
||||
# if "LLM" not in st.session_state:
|
||||
if llm_selected == 'Titan Text Express':
|
||||
LLM = Bedrock(
|
||||
model_id="amazon.titan-text-express-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
"maxTokenCount": 512,
|
||||
"stopSequences": [],
|
||||
"temperature": 0,
|
||||
"topP": 1,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Llama 2 Chat 70B':
|
||||
LLM = Bedrock(
|
||||
model_id="meta.llama2-70b-chat-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
"max_gen_len": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Llama 2 Chat 13B':
|
||||
LLM = Bedrock(
|
||||
model_id="meta.llama2-13b-chat-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
"max_gen_len": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Claude Instant':
|
||||
LLM = Bedrock(
|
||||
model_id="anthropic.claude-instant-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
# "max_tokens_to_sample": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Claude 2':
|
||||
LLM = Bedrock(
|
||||
model_id="anthropic.claude-v2:1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
# "max_tokens_to_sample": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
st.session_state["LLM"] = LLM
|
||||
|
||||
# if "QA" not in st.session_state:
|
||||
# prompt, memory = model_memory()
|
||||
|
||||
# QA = RetrievalQA.from_chain_type(
|
||||
# llm=LLM,
|
||||
# chain_type="stuff",
|
||||
# retriever=RETRIEVER,
|
||||
# return_source_documents=True,
|
||||
# chain_type_kwargs={"prompt": prompt, "memory": memory},
|
||||
# )
|
||||
# st.session_state["QA"] = QA
|
||||
|
||||
# df = pd.DataFrame(columns=['Contract Name','Contract ID','Actual Value Stored','New Extracted value','Confidence Level','Snippet',
|
||||
# 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result'])
|
||||
df = pd.DataFrame(columns=['Contract Name','Raw value','New Extracted value','Confidence Level','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]
|
||||
contract_list = [contract for contract in contract_list if contract.replace(' MU','').replace('_MU','').replace('.txt','') in list(field_values['(internal) Document Name'])]
|
||||
|
||||
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)
|
||||
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]
|
||||
|
||||
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']]
|
||||
|
||||
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')
|
||||
|
||||
# sample = dict(zip(sample.Filename, sample.Answer))
|
||||
# actual_value_list = [sample.get(contract.rsplit('.',1)[0]+'.txt', ' ') for contract in contract_list]
|
||||
|
||||
# AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID')
|
||||
# AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY')
|
||||
# AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN')
|
||||
|
||||
# Setup bedrock
|
||||
bedrock_runtime = boto3.client(
|
||||
service_name="bedrock-runtime",
|
||||
region_name="us-east-1",
|
||||
)
|
||||
|
||||
# Define the retreiver
|
||||
# load the vectorstore
|
||||
if "EMBEDDINGS" not in st.session_state:
|
||||
EMBEDDINGS = BedrockEmbeddings(
|
||||
client=bedrock_runtime,
|
||||
model_id="amazon.titan-embed-text-v1",
|
||||
)
|
||||
st.session_state.EMBEDDINGS = EMBEDDINGS
|
||||
|
||||
if "DB" not in st.session_state:
|
||||
DB = Chroma(
|
||||
persist_directory=PERSIST_DIRECTORY,
|
||||
embedding_function=st.session_state.EMBEDDINGS,
|
||||
client_settings=CHROMA_SETTINGS,
|
||||
)
|
||||
st.session_state.DB = DB
|
||||
|
||||
# if "RETRIEVER" not in st.session_state:
|
||||
# # { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }
|
||||
# RETRIEVER = DB.as_retriever(search_kwargs={"filter": { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }, "k": 2})
|
||||
# st.session_state.RETRIEVER = RETRIEVER
|
||||
|
||||
# if "LLM" not in st.session_state:
|
||||
if llm_selected == 'Titan Text Express':
|
||||
LLM = Bedrock(
|
||||
model_id="amazon.titan-text-express-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
"maxTokenCount": 512,
|
||||
"stopSequences": [],
|
||||
"temperature": 0,
|
||||
"topP": 1,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Llama 2 Chat 70B':
|
||||
LLM = Bedrock(
|
||||
model_id="meta.llama2-70b-chat-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
"max_gen_len": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Llama 2 Chat 13B':
|
||||
LLM = Bedrock(
|
||||
model_id="meta.llama2-13b-chat-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
"max_gen_len": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Claude Instant':
|
||||
LLM = Bedrock(
|
||||
model_id="anthropic.claude-instant-v1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
# "max_tokens_to_sample": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
elif llm_selected == 'Claude 2':
|
||||
LLM = Bedrock(
|
||||
model_id="anthropic.claude-v2:1",
|
||||
client=bedrock_runtime,
|
||||
model_kwargs={
|
||||
# "max_tokens_to_sample": 512,
|
||||
"temperature": 0,
|
||||
# "topP": 0.9,
|
||||
}
|
||||
)
|
||||
st.session_state["LLM"] = LLM
|
||||
|
||||
# if "QA" not in st.session_state:
|
||||
# prompt, memory = model_memory()
|
||||
|
||||
# QA = RetrievalQA.from_chain_type(
|
||||
# llm=LLM,
|
||||
# chain_type="stuff",
|
||||
# retriever=RETRIEVER,
|
||||
# return_source_documents=True,
|
||||
# chain_type_kwargs={"prompt": prompt, "memory": memory},
|
||||
# )
|
||||
# st.session_state["QA"] = QA
|
||||
|
||||
# df = pd.DataFrame(columns=['Contract Name','Contract ID','Actual Value Stored','New Extracted value','Confidence Level','Snippet',
|
||||
# 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result'])
|
||||
df = pd.DataFrame(columns=['Contract Name','Raw value','New Extracted value','Confidence Level','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)
|
||||
history = pd.read_csv('history.csv')
|
||||
except:
|
||||
accuracy = 'NA'
|
||||
history = pd.DataFrame(columns=['Field Name','# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #'])
|
||||
attempt = 0
|
||||
|
||||
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)
|
||||
if llm_selected in ['Llama 2 Chat 13B', 'Llama 2 Chat 70B']:
|
||||
k_value = 10
|
||||
else:
|
||||
k_value = 20
|
||||
|
||||
|
||||
# @st.cache_data
|
||||
# def convert_df(df):
|
||||
# return df.to_csv(index=False).encode('utf-8')
|
||||
if st.button("Test Configuration"):
|
||||
|
||||
# csv = convert_df(edited_df)
|
||||
answer_list = []
|
||||
doc_list = []
|
||||
response_list = []
|
||||
score_list = []
|
||||
attempt = attempt + 1
|
||||
|
||||
# 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")
|
||||
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.write(column_name)
|
||||
st.write(len(contract_list))
|
||||
# @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")
|
||||
|
||||
|
||||
|
||||
Reference in New Issue
Block a user