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
2024-03-27 13:24:33 +05:30
from sf_conn import get_secret , save_to_sf
from io import StringIO
2024-03-14 17:05:39 +05:30
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-04-02 14:11:37 +05:30
, ' vnair@aarete.com ' , ' kminhas@aarete.com ' , ' fmohiuddin@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-04-02 14:31:43 +05:30
_ , c1 = st . columns ( [ 5 , 1 ] )
2024-04-05 12:41:40 +05:30
# util.setup_page(REDIRECT_URI)
2024-03-13 17:23:20 +05:30
try :
2024-04-02 22:18:04 +05:30
util . setup_page ( REDIRECT_URI )
2024-03-13 17:23:20 +05:30
except :
2024-04-05 12:40:26 +05:30
st . write ( " SSO Failed " )
st . session_state [ ' user_info ' ] = { ' mail ' : ' maamseek@aarete.com ' , ' displayName ' : ' Mayank Aamseek ' }
2024-04-05 12:15:54 +05:30
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 ( )
2024-03-13 17:23:20 +05:30
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-28 23:31:12 +05:30
field_values = field_values . loc [ : , ~ field_values . columns . str . contains ( ' Unnamed: ' ) ]
2024-03-27 13:24:33 +05:30
st . write ( " Local copy of TRAINING_DATA_RAW table loaded " )
2024-03-15 13:41:19 +05:30
try :
2024-04-02 17:17:21 +05:30
# error('table is not updated')
2024-03-26 15:24:40 +05:30
query = ' select * from " BUSINESS_CONFIG " '
2024-03-15 13:41:19 +05:30
cur . execute ( query )
2024-03-26 15:39:20 +05:30
fields = pd . DataFrame . from_records ( iter ( cur ) , columns = [ x [ 0 ] for x in cur . description ] )
2024-03-28 11:16:04 +05:30
fields . rename ( columns = { ' FIELD_NAME ' : ' Field Name ' } , inplace = True )
2024-03-26 15:24:40 +05:30
fields . rename ( columns = { ' QUESTION ' : ' Interrogation Question? ' } , inplace = True )
2024-03-28 11:16:04 +05:30
fields . rename ( columns = { ' SF_COL_NAME ' : ' SF_DB_COL_NAME ' } , inplace = True )
2024-03-26 15:57:30 +05:30
fields = fields [ ~ fields [ ' SF_DB_COL_NAME ' ] . str . endswith ( ' _PG ' , na = None ) ]
2024-03-26 15:24:40 +05:30
fields [ ' Field Name ' ] = fields [ ' SF_DB_COL_NAME ' ]
2024-03-26 15:36:53 +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 ' )
2024-03-28 23:31:12 +05:30
fields [ ' Field Name ' ] = fields [ ' SF_DB_COL_NAME ' ]
2024-03-15 13:41:19 +05:30
fields = fields [ ~ fields [ ' Field Name ' ] . isnull ( ) ]
2024-04-02 17:17:21 +05:30
st . write ( " Local copy of BUSINESS_CONFIG table loaded " )
2024-03-15 13:41:19 +05:30
2024-03-27 13:24:33 +05:30
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 )
2024-03-27 13:58:40 +05:30
history = history [ [ ' Field Name ' , ' # Contracts Tested ' , ' Username ' , ' Date/Time ' , ' Accuracy ' , ' Attempt # ' ] ]
2024-03-27 13:24:33 +05:30
except :
history = pd . read_csv ( ' history.csv ' )
st . write ( " Local copy of TRAINING_ATTEMPT_LOGS table loaded " )
2024-03-13 17:23:20 +05:30
2024-03-29 15:56:53 +05:30
if st . session_state . user_info [ ' 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 ] :
2024-04-06 17:38:53 +05:30
field_group = st . selectbox ( ' Field Group ' , ( ' Unique and Contract Related ' , ' Pricing Before Carveouts - I '
2024-03-12 23:14:37 +05:30
, ' 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-04-06 17:38:53 +05:30
if field_group == ' Unique and Contract Related ' :
fields = fields [ fields [ ' PRIORITY ' ] . isin ( [ ' A ' , ' C ' ] ) ]
2024-03-12 23:14:37 +05:30
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 ' ,
2024-04-02 19:20:12 +05:30
region_name = " us-east-2 "
2024-03-18 17:27:08 +05:30
)
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-04-05 13:01:50 +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()]
field_values = field_values [ ~ field_values [ field ] . isnull ( ) ]
2024-03-26 15:24:40 +05:30
# df = pd.DataFrame({'col':contract_list})
# st.write(df)
2024-04-05 12:43:20 +05:30
# st.write(len(contract_list))
2024-04-05 12:15:54 +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'])]
contract_list = [ contract for contract in contract_list if contract . rsplit ( ' / ' , 1 ) [ 1 ] . replace ( ' .txt ' , ' .pdf ' ) in list ( field_values [ ' Document_Name ' ] ) ]
2024-04-05 12:43:20 +05:30
# st.write(len(contract_list))
2024-04-05 12:15:54 +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
2024-03-27 18:59:29 +05:30
# st.button("Save Prompt")
2024-03-08 18:15:55 +05:30
with prompt_row [ 0 ] :
2024-03-27 18:59:29 +05:30
prompt = st . text_area ( " **Prompt** " , sequence_input , height = 100 , label_visibility = " collapsed " )
2024-03-08 18:15:55 +05:30
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-04-06 14:49:24 +05:30
# 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 :
question = json . dumps ( { field : prompt , str ( field ) + ' _PG ' : " On which page can I find answer to the question - " + str ( prompt ) } )
# st.write(question)
2024-03-12 23:14:37 +05:30
question_with_schema = question
attempt = 0
2024-03-20 18:43:29 +05:30
2024-03-27 13:24:33 +05:30
def run_llm ( attempt , bucket , contract_list , llm_selected , field , field_values , history ) :
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 ' ] )
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-04-06 14:49:24 +05:30
# st.write(question_with_schema)
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-04-06 17:38:53 +05:30
" maxTokenCount " : 2048 ,
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-04-06 17:38:53 +05:30
" max_gen_len " : 2048 ,
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-04-06 17:38:53 +05:30
" max_tokens_to_sample " : 2048 ,
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 ( " } " ) }
2024-04-06 14:49:24 +05:30
# 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]
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 ]
2024-03-22 23:59:54 +05:30
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 ""
2024-04-06 14:49:24 +05:30
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 ]
2024-03-22 23:59:54 +05:30
snippet_l = [ ' ' . join ( context [ : location ] . split ( ' . ' ) [ - 4 : ] ) + ' ' + ' ' . join ( context [ location : ] . split ( ' . ' ) [ : 5 ]
) if location != - 1 else ' ' for location in location_l ]
2024-04-06 14:49:24 +05:30
# 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]
2024-03-22 23:59:54 +05:30
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 :
2024-04-08 15:34:26 +05:30
answer_list = [ answer . strip ( " \n " ) . strip ( ) . strip ( " { " ) . strip ( " } " ) . strip ( ' " ' ) . rstrip ( ' " ' ) . strip (
' ' ) if answer is not None else None for answer in answer_list ]
2024-03-22 23:59:54 +05:30
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 ]
2024-04-08 15:34:26 +05:30
answer_list = [ answer if " do not see " not in str ( answer ) else " " for answer in answer_list ]
2024-03-22 23:59:54 +05:30
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 ]
2024-04-08 15:34:26 +05:30
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 " Don ' t know " 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 ]
2024-03-22 23:59:54 +05:30
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-27 15:17:29 +05:30
# df['Contract ID'] = contract_list_f
df [ ' Contract ID ' ] = range ( len ( contract_list_f ) )
2024-03-08 18:15:55 +05:30
# to be deleted later
2024-04-05 12:15:54 +05:30
# contract_list_f = [contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list_f]
2024-04-06 14:49:24 +05:30
contract_list_f = [ contract . rsplit ( ' / ' , 1 ) [ 1 ] . replace ( ' .txt ' , ' .pdf ' ) 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-04-08 15:34:26 +05:30
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 " " )
2024-03-22 23:59:54 +05:30
df [ ' Revised Prompt ' ] = [ prompt ] * len ( contract_list_f )
2024-03-28 23:31:12 +05:30
2024-03-22 23:59:54 +05:30
df = pd . merge ( df , fields [ [ ' Field Name ' , ' SF_DB_COL_NAME ' ] ] , how = ' left ' , on = ' Field Name ' )
2024-03-28 23:31:12 +05:30
field_values_2 = pd . DataFrame ( columns = [ ' Contract Name ' , ' SF_DB_COL_NAME ' , ' Actual Value Stored ' , ' Original Page Number ' ] )
2024-03-22 23:59:54 +05:30
for file_name in contract_list :
2024-04-05 12:15:54 +05:30
# 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 ]
2024-03-22 23:59:54 +05:30
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 ' ]
2024-03-28 23:31:12 +05:30
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 ' ] )
2024-03-22 23:59:54 +05:30
field_values_1 [ ' Contract Name ' ] = document_name
field_values_2 = pd . concat ( [ field_values_2 , field_values_1 ] , ignore_index = True )
2024-03-29 15:36:46 +05:30
2024-03-22 23:59:54 +05:30
df = pd . merge ( df , field_values_2 , how = ' left ' , on = [ ' Contract Name ' , ' SF_DB_COL_NAME ' ] )
2024-04-08 15:34:26 +05:30
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 " " )
2024-03-08 18:15:55 +05:30
2024-04-08 15:34:26 +05:30
df_date = df [ df [ ' SF_DB_COL_NAME ' ] . str . contains ( ' _DT ' ) ]
df_others = df [ ~ df [ ' SF_DB_COL_NAME ' ] . str . contains ( ' _DT ' ) ]
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 )
2024-03-14 22:53:13 +05:30
2024-03-08 18:15:55 +05:30
df . fillna ( " " , inplace = True )
2024-04-08 15:34:26 +05:30
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 ' ' )
2024-03-08 18:15:55 +05:30
actual_value_list = list ( df [ ' Actual Value Stored ' ] )
2024-04-08 15:34:26 +05:30
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 ) ]
2024-03-08 18:15:55 +05:30
df [ ' Result ' ] = [ str ( x ) for x in result_list ]
df = df [ ~ df [ ' Contract ID ' ] . isnull ( ) ]
2024-04-08 15:34:26 +05:30
2024-03-27 18:59:29 +05:30
df = df [ [ ' Contract Name ' , ' Contract ID ' , ' SF_DB_COL_NAME ' , ' Actual Value Stored ' , ' New Extracted value ' , ' Confidence Level '
2024-03-22 23:59:54 +05:30
, ' 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 '
2024-03-27 13:24:33 +05:30
if mode == ' Multiple fields ' :
field = field_group
2024-03-29 15:56:53 +05:30
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 ]
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 = ' '
2024-03-27 13:58:40 +05:30
df = pd . DataFrame ( columns = [ ' Contract Name ' , ' Contract ID ' , ' Actual Value Stored ' , ' New Extracted value ' , ' Confidence Level '
, ' Snippet ' , ' New Page Number ' , ' Revised Prompt ' , ' Result ' ] )
2024-03-27 13:24:33 +05:30
2024-03-22 23:59:54 +05:30
if st . button ( " Test Configuration " ) :
2024-03-27 13:24:33 +05:30
df , history , attempt , raw_response_text = run_llm ( attempt , bucket , contract_list , llm_selected , field , field_values , history )
2024-03-27 18:59:29 +05:30
df_1 = df [ [ ' Contract Name ' , ' Contract ID ' , ' Actual Value Stored ' , ' New Extracted value ' , ' Confidence Level '
2024-03-27 13:24:33 +05:30
, ' Snippet ' , ' Original Page Number ' , ' New Page Number ' , ' Revised Prompt ' , ' Result ' ] ]
2024-04-06 17:38:53 +05:30
df . to_csv ( ' results.csv ' , index = False )
history . to_csv ( ' history.csv ' , index = False )
2024-03-27 13:24:33 +05:30
# s3_client.upload_file('results.csv', bucket, key)
csv_buf = StringIO ( )
2024-03-27 18:59:29 +05:30
df_1 . to_csv ( csv_buf , header = True , index = False )
2024-03-27 13:24:33 +05:30
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 ( )
2024-03-28 21:28:27 +05:30
history . tail ( 1 ) . to_csv ( csv_buf , header = True , index = False )
2024-03-27 13:24:33 +05:30
csv_buf . seek ( 0 )
s3_client . put_object ( Bucket = ' doczy-dev-infra-raw-data-ingestion ' , Body = csv_buf . getvalue ( ) , Key = ' training_interface/history.csv ' )
2024-03-27 18:59:29 +05:30
df = df [ [ ' Contract Name ' , ' Contract ID ' , ' SF_DB_COL_NAME ' , ' Actual Value Stored ' , ' New Extracted value ' , ' Confidence Level '
2024-03-28 23:31:12 +05:30
, ' Snippet ' , ' Original Page Number ' , ' New Page Number ' , ' Revised Prompt ' , ' Result ' ] ]
2024-03-27 13:24:33 +05:30
2024-04-06 17:38:53 +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 )
2024-04-05 13:01:50 +05:30
if mode != ' Multiple fields ' :
st . write ( fields . loc [ fields [ ' SF_DB_COL_NAME ' ] == field , ' Field Name ' ] . iloc [ 0 ] )
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-27 13:24:33 +05:30
try :
2024-04-02 17:11:56 +05:30
save_to_sf ( ' load_training_results ' , " training_results_file_name " , " attempt_logs_file_name " , " results.csv " , " history.csv " )
2024-03-27 13:24:33 +05:30
except :
st . write ( " running locally " )
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