Files
doczyai-pipelines/streamlit/interface_3.py
T

543 lines
24 KiB
Python
Raw Normal View History

2024-02-28 18:23:22 +05:30
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
2024-03-12 23:14:37 +05:30
import numpy as np
2024-02-28 18:23:22 +05:30
from datetime import datetime
import random
import os
import dateutil
2024-03-08 18:15:55 +05:30
import util
2024-03-12 23:14:37 +05:30
import anthropic
import re
2024-03-14 17:05:39 +05:30
import snowflake.connector
from sf_conn import get_secret
2024-02-28 18:23:22 +05:30
2024-03-13 17:23:20 +05:30
REDIRECT_URI = 'https://doczy.aarete.com:8503'
2024-03-18 17:27:08 +05:30
user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
2024-03-08 18:15:55 +05:30
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
2024-03-14 17:05:39 +05:30
, 'vnair@aarete.com']
2024-02-28 18:23:22 +05:30
2024-03-08 18:15:55 +05:30
st.set_page_config(layout = "wide")
2024-02-28 18:23:22 +05:30
# 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")
2024-03-13 17:23:20 +05:30
try:
2024-03-26 15:24:40 +05:30
# util.setup_page(REDIRECT_URI)
2024-03-13 17:23:20 +05:30
_,c1= st.columns([4,1])
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
user_mail = st.session_state.user_info['mail']
except:
user_mail = 'maamseek@aarete.com'
2024-03-15 13:41:19 +05:30
try:
2024-03-15 10:57:38 +00:00
sf_secrets = json.loads(get_secret())
2024-03-15 13:41:19 +05:30
conn = snowflake.connector.connect(
2024-03-15 10:57:38 +00:00
user=sf_secrets.get('user'),
password=sf_secrets.get('password'),
2024-03-15 13:41:19 +05:30
account="aarete-doczyai",
2024-03-15 10:57:38 +00:00
role = "DEVADMIN",
warehouse="DEV_XS",
database="DOCZY_DEV",
2024-03-15 13:41:19 +05:30
schema="STG"
)
cur = conn.cursor()
query = 'select * from "TRAINING_DATA_RAW"'
cur.execute(query)
2024-03-15 10:57:38 +00:00
field_values = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
# st.write(field_values)
2024-03-18 17:27:08 +05:30
field_values['Document_Name'] = field_values['DOCUMENT_NAME']
# field_values['Contract ID'] = field_values['CONTRACT_TITLE']
# error('table values are incorrect')
2024-03-15 13:41:19 +05:30
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)
2024-03-18 17:27:08 +05:30
# field_values.rename(columns={'(Internal) Carveout ID': 'Contract ID'}, inplace = True)
2024-03-15 10:57:38 +00:00
# st.write("conn failed")
2024-03-15 13:41:19 +05:30
try:
2024-03-26 15:24:40 +05:30
query = 'select * from "BUSINESS_CONFIG"'
2024-03-15 13:41:19 +05:30
cur.execute(query)
fields = pd.DataFrame(cur.fetchall())
2024-03-26 15:24:40 +05:30
# fields.rename(columns={'FIELD_NAME': 'Field Name'}, inplace = True)
fields.rename(columns={'QUESTION': 'Interrogation Question?'}, inplace = True)
2024-03-22 23:59:54 +05:30
fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True)
fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True)
2024-03-26 15:24:40 +05:30
fields['Field Name'] = fields['SF_DB_COL_NAME']
2024-03-22 23:59:54 +05:30
fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True)
2024-03-26 15:24:40 +05:30
# error('table is empty')
2024-03-15 13:41:19 +05:30
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()]
2024-03-26 15:24:40 +05:30
st.write("conn failed")
2024-03-15 13:41:19 +05:30
2024-03-13 17:23:20 +05:30
if user_mail in user_list:
2024-03-11 16:05:48 +05:30
2024-03-12 23:14:37 +05:30
field_row = st.columns([0.15, 0.45, 0.4])
with field_row[0]:
st.write("**Field Group**")
with field_row[1]:
field_group = st.selectbox('Field Group',('Unique Key', 'Contract Related', 'Pricing Before Carveouts - I'
, 'Pricing Before Carveouts - II', 'Carveout Indicator, Code Type and Code #s - I'
, 'Carveout Indicator, Code Type and Code #s - II', 'Carveout Indicator, Code Type and Code #s - III'
, 'Optimize Carving Indic.', 'Carveout Method - I', 'Carveout Method - II', 'Provider'
, 'Timeline'), label_visibility = "collapsed")
2024-03-15 13:41:19 +05:30
# priorty column will be relaced by group_id in snowflake db
2024-03-12 23:14:37 +05:30
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(' ')
2024-03-08 18:15:55 +05:30
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
2024-02-28 18:23:22 +05:30
2024-03-22 23:59:54 +05:30
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', 'Single field - Only Non Empty values', 'Multiple fields'), index=0, label_visibility = "collapsed")
2024-03-08 18:15:55 +05:30
field_row = st.columns([0.15, 0.45, 0.4])
with field_row[0]:
2024-03-22 23:59:54 +05:30
if mode != 'Multiple fields':
st.write("**Field Name**")
2024-03-08 18:15:55 +05:30
with field_row[1]:
2024-03-22 23:59:54 +05:30
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()))
2024-03-05 22:28:13 +05:30
2024-03-08 18:15:55 +05:30
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]:
2024-03-18 17:27:08 +05:30
contract_count = st.selectbox('Contract count',('1', '10', '20', '30', '50', '100', '200', 'All'), index=1, label_visibility = "collapsed")
2024-03-05 22:28:13 +05:30
2024-03-18 17:27:08 +05:30
s3_client = boto3.client('s3',
region_name="us-east-1"
)
bucket = 'doczy-dev-infra-textract'
2024-03-20 18:43:29 +05:30
objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/contract-text-file/")
2024-03-18 17:27:08 +05:30
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))
2024-03-05 22:28:13 +05:30
2024-03-22 23:59:54 +05:30
if mode == 'Single field - Only 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()]
2024-03-05 20:20:22 +05:30
# to be deleted later
2024-03-26 15:24:40 +05:30
# df = pd.DataFrame({'col':contract_list})
# st.write(df)
2024-03-18 17:27:08 +05:30
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'])]
2024-03-26 15:24:40 +05:30
2024-03-18 17:27:08 +05:30
if contract_count == 'All':
contract_count = len(contract_list)
2024-03-22 23:59:54 +05:30
seed_row = st.columns([0.15, 0.45, 0.4])
2024-03-08 18:15:55 +05:30
with seed_row[0]:
2024-03-18 17:27:08 +05:30
if contract_count in ['10', '20', '30', '50', '100', '200']:
2024-03-08 18:15:55 +05:30
st.write("**Seed Value**")
elif contract_count == '1':
st.write("**Contract Name**")
with seed_row[1]:
2024-03-18 17:27:08 +05:30
if contract_count in ['10', '20', '30', '50', '100', '200']:
2024-03-08 18:15:55 +05:30
seed_value = st.text_input("**Seed Value**", value = 20, label_visibility = "collapsed")
random.seed(seed_value)
2024-03-18 17:27:08 +05:30
# contract_list = sorted(random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count)))
2024-03-22 23:59:54 +05:30
contract_list = sorted(random.choices(contract_list, k=int(contract_count)))
2024-03-08 18:15:55 +05:30
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]:
2024-03-12 23:14:37 +05:30
llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude Instant', 'Llama 2 Chat 70B'
, 'Titan Text Express'), index=1, label_visibility = "collapsed")
2024-03-08 18:15:55 +05:30
st.write("**Prompt**")
2024-03-22 23:59:54 +05:30
if mode == 'Multiple fields':
sequence_input = json.dumps(field_prompt_mapping)
else:
sequence_input = field_prompt_mapping.get(field)
2024-03-08 18:15:55 +05:30
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")
2024-03-22 23:59:54 +05:30
# 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)
2024-03-08 18:15:55 +05:30
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",
)
2024-03-12 23:14:37 +05:30
question = prompt
question_with_schema = question
attempt = 0
2024-03-20 18:43:29 +05:30
2024-03-22 23:59:54 +05:30
def run_llm(attempt, bucket, contract_list, llm_selected, field_values):
2024-03-08 18:15:55 +05:30
2024-03-22 23:59:54 +05:30
# 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'])
try:
history = pd.read_csv('history.csv')
except:
history = pd.DataFrame(columns=['Field Name','# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #'])
2024-03-08 18:15:55 +05:30
2024-03-22 23:59:54 +05:30
field_list = []
2024-03-08 18:15:55 +05:30
answer_list = []
2024-03-12 23:14:37 +05:30
snippet_list = []
page_no_list = []
2024-03-22 23:59:54 +05:30
contract_list_f = []
2024-03-08 18:15:55 +05:30
attempt = attempt + 1
2024-03-12 23:14:37 +05:30
for contract in contract_list:
2024-03-18 17:27:08 +05:30
# 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")
2024-03-12 23:14:37 +05:30
# Add "You must answer in correct JSON format."
# Add Answer in JSON format: {{
if llm_selected == "Titan Text Express":
context = context[:16000]
2024-03-22 23:59:54 +05:30
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.
2024-03-12 23:14:37 +05:30
#
{context}
#
Question: {question}
2024-03-22 23:59:54 +05:30
Answer: Answer in JSON format: {{"""
2024-03-12 23:14:37 +05:30
parameters = {
2024-03-22 23:59:54 +05:30
"maxTokenCount":1024,
2024-03-12 23:14:37 +05:30
"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':
2024-03-20 18:43:29 +05:30
context = context[:6000]
2024-03-22 23:59:54 +05:30
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.
2024-03-12 23:14:37 +05:30
##
{context}
##
Question: {question}
2024-03-22 23:59:54 +05:30
Answer: Answer in JSON format: {{"""
2024-03-12 23:14:37 +05:30
payload={
"prompt":"[INST]"+ prompt_data +"[/INST]",
2024-03-22 23:59:54 +05:30
"max_gen_len":1024,
2024-03-12 23:14:37 +05:30
"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']:
2024-03-20 18:43:29 +05:30
if llm_selected == 'Claude Instant':
2024-03-22 23:59:54 +05:30
context = context[:175000]
2024-03-12 23:14:37 +05:30
prompt_data = f"""
2024-03-22 23:59:54 +05:30
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.
2024-03-12 23:14:37 +05:30
{context}
Question: {question_with_schema}
2024-03-22 23:59:54 +05:30
Assistant: Answer in JSON format: {{"""
2024-03-12 23:14:37 +05:30
body = json.dumps(
{"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT,
2024-03-22 23:59:54 +05:30
"max_tokens_to_sample": 1024,
2024-03-12 23:14:37 +05:30
"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"
2024-03-20 18:43:29 +05:30
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']
2024-03-22 23:59:54 +05:30
2024-03-20 18:43:29 +05:30
except:
response_text = "failed"
2024-03-12 23:14:37 +05:30
2024-03-22 23:59:54 +05:30
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 + "}"
2024-03-12 23:14:37 +05:30
try:
response_dict = json.loads(response_text)
except:
2024-03-22 23:59:54 +05:30
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, " ")
2024-03-08 18:15:55 +05:30
try:
2024-03-22 23:59:54 +05:30
if isinstance(response_dict, dict):
answer_l = list(response_dict.values())[:1]
else:
answer_l = list(response_dict)[:1]
2024-03-08 18:15:55 +05:30
except:
2024-03-22 23:59:54 +05:30
answer_l = [response_dict]
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 ""
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
try:
answer_list = [answer.strip("\n").strip().strip("{").strip("}").strip('"').rstrip('"') 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 "does not specify" not in str(answer) else " " for answer in answer_list]
answer_list = [answer if "does not explicitly" not in str(answer) else " " for answer in answer_list]
answer_list = [answer.rstrip(".") for answer in answer_list]
else:
answer_list = [answer.rstrip(".") for answer in answer_list]
# answer_list = [str(x).rsplit(':',1)[0] if len(str(x).rsplit(':',1)) < 2 else str(x).rsplit(':',1)[1] for x in answer_list]
except:
print('post processing failed')
2024-03-14 22:53:13 +05:30
2024-03-22 23:59:54 +05:30
df['Contract ID'] = contract_list_f
2024-03-08 18:15:55 +05:30
# to be deleted later
2024-03-22 23:59:54 +05:30
contract_list_f = [contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list_f]
2024-03-14 22:53:13 +05:30
2024-03-22 23:59:54 +05:30
df['Contract Name'] = contract_list_f
2024-03-08 18:15:55 +05:30
df['New Extracted value'] = answer_list
2024-03-12 23:14:37 +05:30
df['Confidence Level'] = ' '
df['Snippet'] = snippet_list
df['New Page Number'] = page_no_list
2024-03-22 23:59:54 +05:30
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'])
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]
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_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'])
2024-03-08 18:15:55 +05:30
answer_list = list(df['New Extracted value'])
2024-03-22 23:59:54 +05:30
# if 'Date' in field:
# df['Actual Value Stored'] = pd.to_datetime(df['Actual Value Stored'],errors='coerce').dt.date
2024-03-14 22:53:13 +05:30
2024-03-08 18:15:55 +05:30
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()]
2024-03-22 23:59:54 +05:30
if 'Original Page Number' not in df.columns:
df['Original Page Number'] = ' '
df = df[['Contract Name','Contract ID', 'Field Name', 'SF_DB_COL_NAME', 'Actual Value Stored','New Extracted value','Confidence Level'
,'Snippet','Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
2024-03-08 18:15:55 +05:30
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]
2024-03-11 16:05:48 +05:30
# df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False)
2024-03-22 23:59:54 +05:30
return df, history, attempt, raw_response_text
raw_response_text = ''
if st.button("Test Configuration"):
df, history, attempt, raw_response_text = run_llm(attempt, bucket, contract_list, llm_selected, field_values)
df.to_csv('results.csv', index=False)
2024-03-08 18:15:55 +05:30
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()
2024-03-22 23:59:54 +05:30
df = pd.read_csv('results.csv')
df['Result'] = df['Result'].astype('str')
history = pd.read_csv('history.csv')
2024-03-08 18:15:55 +05:30
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")
2024-03-22 23:59:54 +05:30
add_vertical_space(20)
st.write(field)
# st.write(column_name)
2024-03-08 18:15:55 +05:30
st.write(len(contract_list))
2024-03-22 23:59:54 +05:30
st.write(raw_response_text)
2024-02-28 18:23:22 +05:30
2024-03-08 18:15:55 +05:30
else:
st.write("Access Denied")
2024-03-05 20:20:22 +05:30
2024-02-28 18:23:22 +05:30