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