Files
doczyai-pipelines/streamlit/interface_3_rag.py
T

537 lines
18 KiB
Python
Raw Normal View History

2024-03-12 23:14:37 +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
2024-09-30 12:21:29 +01:00
from constants import (
CHROMA_SETTINGS,
EMBEDDING_MODEL_NAME,
PERSIST_DIRECTORY,
MODEL_ID,
MODEL_BASENAME,
SOURCE_DIRECTORY,
USER_LIST,
)
2024-03-12 23:14:37 +05:30
from langchain.chains import RetrievalQA
import streamlit as st
from streamlit_extras.add_vertical_space import add_vertical_space
import pandas as pd
from datetime import datetime
import random
import os
import dateutil
import util
2024-09-30 12:21:29 +01:00
REDIRECT_URI = "http://172.29.20.126:8503"
2024-05-06 19:28:00 -05:00
user_list = USER_LIST
2024-03-12 23:14:37 +05:30
2024-09-30 12:21:29 +01:00
st.set_page_config(layout="wide")
2024-03-12 23:14:37 +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")
# util.setup_page(REDIRECT_URI)
2024-09-30 12:21:29 +01:00
# if st.session_state.user_info['mail'] in user_list:
if "maamseek@aarete.com" in user_list:
2024-03-12 23:14:37 +05:30
2024-09-30 12:21:29 +01:00
fields = pd.read_csv(
"contract_fields.csv", encoding="unicode_escape", skipinitialspace=True
)
fields = fields[fields["PRIORITY"].isin(["A", "C"])]
field_values = pd.read_csv(
"contract_field_values.csv", encoding="unicode_escape", skipinitialspace=True
)
fields = fields.drop_duplicates(subset="Field Name", keep="first").sort_values(
"Field Name"
)
fields = fields[~fields["Field Name"].isnull()]
field_prompt_mapping = dict(
zip(fields["Field Name"], fields["Interrogation Question?"])
)
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 Name**")
with field_row[1]:
2024-09-30 12:21:29 +01:00
field = st.selectbox(
"Field Name",
sorted(set(field_prompt_mapping.keys())),
index=0,
label_visibility="collapsed",
)
2024-03-12 23:14:37 +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-09-30 12:21:29 +01:00
contract_count = st.selectbox(
"Contract count",
("1", "10", "20", "30", "50", "All"),
index=1,
label_visibility="collapsed",
)
2024-03-12 23:14:37 +05:30
seed_row = st.columns([0.15, 0.45, 0.4])
contract_list = sorted(os.listdir(SOURCE_DIRECTORY))
# to be deleted later
2024-09-30 12:21:29 +01:00
contract_list = [
contract
for contract in contract_list
if contract.replace(" MU", "").replace("_MU", "").replace(".txt", "")
in list(field_values["(internal) Document Name"])
]
2024-03-12 23:14:37 +05:30
with seed_row[0]:
2024-09-30 12:21:29 +01:00
if contract_count in ["10", "20", "30", "50"]:
2024-03-12 23:14:37 +05:30
st.write("**Seed Value**")
2024-09-30 12:21:29 +01:00
elif contract_count == "1":
2024-03-12 23:14:37 +05:30
st.write("**Contract Name**")
with seed_row[1]:
2024-09-30 12:21:29 +01:00
if contract_count in ["10", "20", "30", "50"]:
seed_value = st.text_input(
"**Seed Value**", value=20, label_visibility="collapsed"
)
2024-03-12 23:14:37 +05:30
random.seed(seed_value)
2024-09-30 12:21:29 +01:00
contract_list = sorted(
random.choices(os.listdir(SOURCE_DIRECTORY), k=int(contract_count))
)
elif contract_count == "1":
contract_name = st.selectbox(
"Contract Name", (contract_list), label_visibility="collapsed"
)
2024-03-12 23:14:37 +05:30
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-09-30 12:21:29 +01:00
llm_selected = st.selectbox(
"Langauge Model",
(
"Claude 2",
"Claude Instant",
"Llama 2 Chat 13B",
"Llama 2 Chat 70B",
"Titan Text Express",
),
label_visibility="collapsed",
)
2024-03-12 23:14:37 +05:30
st.write("**Prompt**")
sequence_input = field_prompt_mapping.get(field)
prompt_row = st.columns([0.8, 0.2])
with prompt_row[1]:
if st.button("Clear Prompt"):
2024-09-30 12:21:29 +01:00
sequence_input = ""
2024-03-12 23:14:37 +05:30
if st.button("Back to default"):
prompt = sequence_input
st.button("Save Prompt")
with prompt_row[0]:
2024-09-30 12:21:29 +01:00
prompt = st.text_area(
"**Prompt**", sequence_input, height=150, label_visibility="collapsed"
)
2024-03-12 23:14:37 +05:30
page_list_all = []
for contract in contract_list:
page_list = []
2024-09-30 12:21:29 +01:00
with open(
os.path.join(SOURCE_DIRECTORY, contract[:-4] + ".txt"), "r"
) as infile:
2024-03-12 23:14:37 +05:30
text = infile.read()
2024-09-30 12:21:29 +01:00
page_count = text.count("Start of Page No. = ")
for page in range(page_count + 1):
file_path = "SOURCE_DOCUMENTS\\" + f"{contract[:-4]}_page{page}.txt"
dict_with_pages = {"source": {"$eq": file_path}}
2024-03-12 23:14:37 +05:30
page_list.append(dict_with_pages)
page_list_all.append(page_list)
contract_txt_mapping = dict(zip(contract_list, page_list_all))
2024-09-30 12:21:29 +01:00
column_name = fields.loc[fields["Field Name"] == field, "SF_DB_COL_NAME"].iloc[0]
column_list = ["(internal) Document Name", "(Internal) Carveout ID", column_name]
if column_name + "_PG" in list(field_values.columns):
column_list.append(column_name + "_PG")
2024-03-12 23:14:37 +05:30
field_values = field_values[column_list]
2024-09-30 12:21:29 +01:00
field_values.rename(
columns={
"(internal) Document Name": "Contract Name",
column_name: "Actual Value Stored",
"(Internal) Carveout ID": "Contract ID",
column_name + "_PG": "Original Page Number",
},
inplace=True,
)
field_values = field_values.drop_duplicates(
subset="Contract Name", keep="first"
).sort_values("Contract Name")
2024-03-12 23:14:37 +05:30
# Setup bedrock
bedrock_runtime = boto3.client(
service_name="bedrock-runtime",
region_name="us-east-1",
)
# Define the retreiver
# load the vectorstore
if "EMBEDDINGS" not in st.session_state:
EMBEDDINGS = BedrockEmbeddings(
client=bedrock_runtime,
model_id="amazon.titan-embed-text-v1",
)
st.session_state.EMBEDDINGS = EMBEDDINGS
if "DB" not in st.session_state:
DB = Chroma(
persist_directory=PERSIST_DIRECTORY,
embedding_function=st.session_state.EMBEDDINGS,
client_settings=CHROMA_SETTINGS,
)
st.session_state.DB = DB
# if "RETRIEVER" not in st.session_state:
# # { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }
# RETRIEVER = DB.as_retriever(search_kwargs={"filter": { "source": { '$eq': "SOURCE_DOCUMENTS\\A.1_UH_Health_System_eff_2_1_08 (1)_page0.txt"} }, "k": 2})
# st.session_state.RETRIEVER = RETRIEVER
# if "LLM" not in st.session_state:
2024-09-30 12:21:29 +01:00
if llm_selected == "Titan Text Express":
2024-03-12 23:14:37 +05:30
LLM = Bedrock(
model_id="amazon.titan-text-express-v1",
client=bedrock_runtime,
model_kwargs={
"maxTokenCount": 512,
"stopSequences": [],
"temperature": 0,
"topP": 1,
2024-09-30 12:21:29 +01:00
},
2024-03-12 23:14:37 +05:30
)
2024-09-30 12:21:29 +01:00
elif llm_selected == "Llama 2 Chat 70B":
2024-03-12 23:14:37 +05:30
LLM = Bedrock(
model_id="meta.llama2-70b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
2024-09-30 12:21:29 +01:00
},
2024-03-12 23:14:37 +05:30
)
2024-09-30 12:21:29 +01:00
elif llm_selected == "Llama 2 Chat 13B":
2024-03-12 23:14:37 +05:30
LLM = Bedrock(
model_id="meta.llama2-13b-chat-v1",
client=bedrock_runtime,
model_kwargs={
"max_gen_len": 512,
"temperature": 0,
# "topP": 0.9,
2024-09-30 12:21:29 +01:00
},
2024-03-12 23:14:37 +05:30
)
2024-09-30 12:21:29 +01:00
elif llm_selected == "Claude Instant":
2024-03-12 23:14:37 +05:30
LLM = Bedrock(
model_id="anthropic.claude-instant-v1",
client=bedrock_runtime,
model_kwargs={
# "max_tokens_to_sample": 512,
"temperature": 0,
# "topP": 0.9,
2024-09-30 12:21:29 +01:00
},
2024-03-12 23:14:37 +05:30
)
2024-09-30 12:21:29 +01:00
elif llm_selected == "Claude 2":
2024-03-12 23:14:37 +05:30
LLM = Bedrock(
model_id="anthropic.claude-v2:1",
client=bedrock_runtime,
model_kwargs={
# "max_tokens_to_sample": 512,
"temperature": 0,
# "topP": 0.9,
2024-09-30 12:21:29 +01:00
},
2024-03-12 23:14:37 +05:30
)
st.session_state["LLM"] = LLM
# if "QA" not in st.session_state:
# prompt, memory = model_memory()
# QA = RetrievalQA.from_chain_type(
# llm=LLM,
# chain_type="stuff",
# retriever=RETRIEVER,
# return_source_documents=True,
# chain_type_kwargs={"prompt": prompt, "memory": memory},
# )
# st.session_state["QA"] = QA
# df = pd.DataFrame(columns=['Contract Name','Contract ID','Actual Value Stored','New Extracted value','Confidence Level','Snippet',
# 'Original Page Number', 'New Page Number', 'Revised Prompt', 'Result'])
2024-09-30 12:21:29 +01:00
df = pd.DataFrame(
columns=[
"Contract Name",
"Raw value",
"New Extracted value",
"Confidence Level",
"Snippet1",
"Snippet2",
"Snippet3",
"Snippet4",
"Snippet5",
"Snippet6",
"Snippet7",
"Snippet8",
"Snippet9",
"Snippet10",
"New Page Number",
"Revised Prompt",
"Result",
]
)
2024-03-12 23:14:37 +05:30
try:
2024-09-30 12:21:29 +01:00
history = pd.read_csv("history.csv")
2024-03-12 23:14:37 +05:30
except:
2024-09-30 12:21:29 +01:00
history = pd.DataFrame(
columns=[
"Field Name",
"# Contracts Tested",
"Username",
"Date/Time",
"Accuracy",
"Attempt #",
]
)
2024-03-12 23:14:37 +05:30
attempt = 0
2024-09-30 12:21:29 +01:00
if llm_selected in ["Llama 2 Chat 13B", "Llama 2 Chat 70B"]:
2024-03-12 23:14:37 +05:30
k_value = 10
else:
k_value = 20
if st.button("Test Configuration"):
answer_list = []
doc_list = []
response_list = []
score_list = []
attempt = attempt + 1
for page_list in page_list_all:
2024-09-30 12:21:29 +01:00
RETRIEVER = st.session_state.DB.as_retriever(
search_kwargs={"filter": {"$or": page_list}, "k": k_value}
)
2024-03-12 23:14:37 +05:30
QA = RetrievalQA.from_chain_type(
llm=st.session_state["LLM"],
chain_type="stuff",
retriever=RETRIEVER,
return_source_documents=True,
# chain_type_kwargs={"prompt": prompt, "memory": None},
)
2024-09-30 12:21:29 +01:00
score = st.session_state.DB.similarity_search_with_relevance_scores(
prompt, k=4, filter={"$or": page_list}
)
2024-03-12 23:14:37 +05:30
score_list.append(max(d[1] for d in score))
response = QA(prompt)
answer, docs = response["result"], response["source_documents"]
answer_list.append(answer)
doc_list.append(docs)
response_list.append(response)
2024-09-30 12:21:29 +01:00
df["Raw value"] = answer_list
2024-03-12 23:14:37 +05:30
# post-processing
2024-09-30 12:21:29 +01:00
if "Date" in field:
2024-03-12 23:14:37 +05:30
date_list = []
for answer in answer_list:
try:
2024-09-30 12:21:29 +01:00
extracted_date = dateutil.parser.parse(
str(answer).replace('"', ""), fuzzy=True
).date()
2024-03-12 23:14:37 +05:30
except:
extracted_date = " "
date_list.append(extracted_date)
answer_list = date_list
2024-09-30 12:21:29 +01:00
elif llm_selected in ["Llama 2 Chat 13B", "Llama 2 Chat 70B"]:
2024-03-12 23:14:37 +05:30
answer_list = [answer.rstrip(".") for answer in answer_list]
2024-09-30 12:21:29 +01:00
answer_list = [
answer if "I don't know" not in str(answer) else " "
for answer in answer_list
]
answer_list = [
answer if "N/A" not in str(answer) else " " for answer in answer_list
]
answer_list = [
answer if "does not contain" not in str(answer) else " "
for answer in answer_list
]
answer_list = [
answer if "None" not in str(answer) else " " for answer in answer_list
]
answer_list = [
answer if "Not specified in the contract" not in str(answer) else " "
for answer in answer_list
]
answer_list = [
answer if "Not applicable" not in str(answer) else " "
for answer in answer_list
]
elif llm_selected in ["Claude 2", "Claude Instant"]:
answer_list = [
(
answer
if "Unfortunately, I do not have enough context" not in str(answer)
else " "
)
for answer in answer_list
]
2024-03-12 23:14:37 +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]
2024-09-30 12:21:29 +01:00
df["Contract Name"] = contract_list
2024-03-12 23:14:37 +05:30
# to be deleted later
2024-09-30 12:21:29 +01:00
df["Contract Name"] = [
contract.replace(" MU", "").replace("_MU", "").replace(".txt", "")
for contract in contract_list
]
2024-03-12 23:14:37 +05:30
2024-09-30 12:21:29 +01:00
df["New Extracted value"] = answer_list
df["Confidence Level"] = [round(score, 2) for score in score_list]
2024-03-12 23:14:37 +05:30
Snippet = []
count = 0
for i in range(int(contract_count)):
2024-09-30 12:21:29 +01:00
for j in range(10):
2024-03-12 23:14:37 +05:30
try:
content = str(doc_list[i][j].page_content)
except:
content = " "
Snippet.append(content)
# df['Snippet1'] = [str(doc[0].page_content) for doc in doc_list]
2024-09-30 12:21:29 +01:00
df["Snippet1"] = Snippet[: int(contract_count)]
df["Snippet2"] = Snippet[int(contract_count) : 2 * int(contract_count)]
df["Snippet3"] = Snippet[2 * int(contract_count) : 3 * int(contract_count)]
df["Snippet4"] = Snippet[3 * int(contract_count) : 4 * int(contract_count)]
df["Snippet5"] = Snippet[4 * int(contract_count) : 5 * int(contract_count)]
df["Snippet6"] = Snippet[5 * int(contract_count) : 6 * int(contract_count)]
df["Snippet7"] = Snippet[6 * int(contract_count) : 7 * int(contract_count)]
df["Snippet8"] = Snippet[7 * int(contract_count) : 8 * int(contract_count)]
df["Snippet9"] = Snippet[8 * int(contract_count) : 9 * int(contract_count)]
df["Snippet10"] = Snippet[9 * int(contract_count) :]
df["New Page Number"] = [
int(str(doc[0].metadata["source"]).rsplit("_page")[1].replace(".txt", ""))
+ 1
for doc in doc_list
]
df["Revised Prompt"] = [prompt] * len(contract_list)
df = pd.merge(df, field_values, how="left", on="Contract Name")
answer_list = list(df["New Extracted value"])
df["Actual Value Stored"] = pd.to_datetime(
df["Actual Value Stored"], errors="coerce"
).dt.date
2024-03-12 23:14:37 +05:30
df.fillna(" ", inplace=True)
2024-09-30 12:21:29 +01:00
actual_value_list = list(df["Actual Value Stored"])
result_list = [i == j for i, j in zip(actual_value_list, answer_list)]
df["Result"] = [str(x) for x in result_list]
df = df[~df["Contract ID"].isnull()]
if "Original Page Number" in df.columns:
df = df[
[
"Contract Name",
"Contract ID",
"Actual Value Stored",
"Raw value",
"New Extracted value",
"Confidence Level",
"Snippet1",
"Snippet2",
"Snippet3",
"Snippet4",
"Snippet5",
"Snippet6",
"Snippet7",
"Snippet8",
"Snippet9",
"Snippet10",
"Original Page Number",
"New Page Number",
"Revised Prompt",
"Result",
]
]
2024-03-12 23:14:37 +05:30
else:
2024-09-30 12:21:29 +01:00
df = df[
[
"Contract Name",
"Contract ID",
"Actual Value Stored",
"Raw value",
"New Extracted value",
"Confidence Level",
"Snippet1",
"Snippet2",
"Snippet3",
"Snippet4",
"Snippet5",
"Snippet6",
"Snippet7",
"Snippet8",
"Snippet9",
"Snippet10",
"New Page Number",
"Revised Prompt",
"Result",
]
]
2024-03-12 23:14:37 +05:30
try:
2024-09-30 12:21:29 +01:00
accuracy = round(
sum(bool(x) for x in result_list) * 100 / len(list(df["Result"])), 2
)
2024-03-12 23:14:37 +05:30
except:
2024-09-30 12:21:29 +01:00
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-12 23:14:37 +05:30
# df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False)
2024-09-30 12:21:29 +01:00
history.to_csv("history.csv", index=False)
2024-03-12 23:14:37 +05:30
# df_copy = df.set_index(df.columns[0]).copy()
# df_2_copy = history.set_index(history.columns[0]).copy()
st.dataframe(df)
st.dataframe(history)
2024-09-30 12:21:29 +01:00
# @st.cache_data
2024-03-12 23:14:37 +05:30
# 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]:
2024-09-30 12:21:29 +01:00
# st.button("Save All Imputations")
2024-03-12 23:14:37 +05:30
# with buttons[1]:
# st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
# with buttons[2]:
2024-09-30 12:21:29 +01:00
# st.button("Kickoff Database Integration")
2024-03-12 23:14:37 +05:30
st.write(column_name)
st.write(len(contract_list))
else:
st.write("Access Denied")