208 lines
7.1 KiB
Python
208 lines
7.1 KiB
Python
import json
|
|
|
|
import boto3
|
|
from langchain.prompts import PromptTemplate
|
|
from langchain.embeddings.bedrock import BedrockEmbeddings
|
|
from langchain.llms.bedrock import Bedrock
|
|
from langchain_community.vectorstores import Chroma
|
|
from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY
|
|
from langchain.chains import RetrievalQA
|
|
|
|
import streamlit as st
|
|
from streamlit_extras.add_vertical_space import add_vertical_space
|
|
import os
|
|
import pandas as pd
|
|
import util
|
|
|
|
REDIRECT_URI = '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', '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 ™")
|
|
st.markdown(
|
|
"""
|
|
## About
|
|
This app extracts data from contracts
|
|
|
|
"""
|
|
)
|
|
add_vertical_space(15)
|
|
# st.write("Doczy")
|
|
|
|
# util.setup_page(REDIRECT_URI)
|
|
# if st.session_state.user_info['mail'] in user_list:
|
|
if 'maamseek@aarete.com' in user_list:
|
|
|
|
fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
|
|
fields = fields[fields['PRIORITY'] == 'A']
|
|
fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
|
|
fields['Interrogation Question?'].fillna(' ', inplace=True)
|
|
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
|
|
|
|
def file_selector(folder_path='RAW_DOCUMENTS'):
|
|
filenames = os.listdir(folder_path)
|
|
selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed")
|
|
# return os.path.join(folder_path, selected_filename)
|
|
return selected_filename
|
|
|
|
file_row = st.columns([0.2, 0.7, 0.1])
|
|
with file_row[0]:
|
|
st.write("**Contract Name**")
|
|
with file_row[1]:
|
|
# file_name = st.text_input("**Contract Name**", label_visibility = "collapsed")
|
|
file_name = file_selector()
|
|
|
|
# lob_row = st.columns([0.2, 0.7, 0.1])
|
|
# with lob_row[0]:
|
|
# st.write("**LOB**")
|
|
# with lob_row[1]:
|
|
# lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed")
|
|
|
|
llm_row = st.columns([0.2, 0.7, 0.1])
|
|
with llm_row[0]:
|
|
st.write("**Langauge Model**")
|
|
with llm_row[1]:
|
|
llm_selected = st.selectbox('Langauge Model',('Llama 2 Chat 13B', 'Llama 2 Chat 70B', 'Titan Text Express'), label_visibility = "collapsed")
|
|
|
|
page_list = []
|
|
with open(os.path.join(SOURCE_DIRECTORY, file_name[:-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'{file_name[:-4]}_page{page}.txt'
|
|
dict_with_pages = { 'source': { '$eq': file_path }}
|
|
page_list.append(dict_with_pages)
|
|
|
|
# AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID')
|
|
# AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY')
|
|
# AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN')
|
|
|
|
# Setup bedrock
|
|
bedrock_runtime = boto3.client(
|
|
service_name="bedrock-runtime",
|
|
region_name="us-east-1"
|
|
)
|
|
|
|
embeddings = BedrockEmbeddings(
|
|
client=bedrock_runtime,
|
|
model_id="amazon.titan-embed-text-v1",
|
|
)
|
|
DB = Chroma(
|
|
persist_directory=PERSIST_DIRECTORY,
|
|
embedding_function=embeddings,
|
|
client_settings=CHROMA_SETTINGS,
|
|
)
|
|
RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4})
|
|
|
|
if llm_selected == 'Titan Text Express':
|
|
LLM = Bedrock(
|
|
model_id="amazon.titan-text-express-v1",
|
|
client=bedrock_runtime,
|
|
model_kwargs={
|
|
"maxTokenCount": 4096,
|
|
"stopSequences": [],
|
|
"temperature": 0,
|
|
"topP": 1,
|
|
}
|
|
)
|
|
elif llm_selected == 'Llama 2 Chat 70B':
|
|
LLM = Bedrock(
|
|
model_id="meta.llama2-70b-chat-v1",
|
|
client=bedrock_runtime,
|
|
model_kwargs={
|
|
"max_gen_len": 512,
|
|
"temperature": 0,
|
|
# "topP": 0.9,
|
|
}
|
|
)
|
|
else:
|
|
LLM = Bedrock(
|
|
model_id="meta.llama2-13b-chat-v1",
|
|
client=bedrock_runtime,
|
|
model_kwargs={
|
|
"max_gen_len": 512,
|
|
"temperature": 0,
|
|
# "topP": 0.9,
|
|
}
|
|
)
|
|
|
|
template = """
|
|
|
|
Use the following pieces of context to answer the question at the end. If you don't know the answer,\
|
|
just say that you don't know, don't try to make up an answer.
|
|
|
|
{context}
|
|
|
|
Question: {question}
|
|
Answer:"""
|
|
prompt = PromptTemplate(input_variables=["context", "question"], template=template)
|
|
|
|
QA = RetrievalQA.from_chain_type(
|
|
llm=LLM,
|
|
chain_type="stuff",
|
|
retriever=RETRIEVER,
|
|
return_source_documents=True,
|
|
chain_type_kwargs={"prompt": prompt},
|
|
)
|
|
|
|
# query = "In which state or states is the Contract applicable? Answer in one or two words. State name: "
|
|
# response = QA({"query":query})
|
|
# st.write(query)
|
|
# st.write(response['result'])
|
|
# st.write("-----------")
|
|
# st.write(response)
|
|
|
|
# clicked = st.button("Show Results")
|
|
df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level',
|
|
'Field Extracted Value','Imputed Value'])
|
|
field_list = list(field_prompt_mapping.keys())
|
|
query_list = [field_prompt_mapping[x] for x in field_list]
|
|
score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list]
|
|
confidence_list = []
|
|
for score in score_list:
|
|
confidence_list.append(max(d[1] for d in score))
|
|
# st.write(confidence_list)
|
|
|
|
if st.button("Show Results"):
|
|
response_list = [QA({"query":query}) for query in query_list]
|
|
answer_list = [response['result'] for response in response_list]
|
|
doc_list = [response['source_documents'] for response in response_list]
|
|
snippet_list = [str(doc[0].page_content) for doc in doc_list]
|
|
page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list]
|
|
|
|
df['Field Name'] = field_list
|
|
df['Contract Name'] = file_name
|
|
df['Snippet'] = snippet_list
|
|
df['Page Number'] = page_no_list
|
|
df['Confidence Level'] = confidence_list
|
|
df['Field Extracted Value'] = answer_list
|
|
df.to_csv('temp2.csv', index=False)
|
|
|
|
df2 = pd.read_csv('temp2.csv')
|
|
df2['Imputed Value'] = ''
|
|
edited_df = st.data_editor(df2)
|
|
|
|
@st.cache_data
|
|
def convert_df(df):
|
|
return df.to_csv(index=False).encode('utf-8')
|
|
|
|
csv = convert_df(edited_df)
|
|
|
|
buttons = st.columns(3)
|
|
with buttons[0]:
|
|
st.button("Save All Imputations")
|
|
with buttons[1]:
|
|
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
|
|
with buttons[2]:
|
|
st.button("Kickoff Database Integration")
|
|
|
|
else:
|
|
st.write("Access Denied")
|
|
|
|
|
|
|