afb6d5185d
Feature/lesser table caching refactor hybrid * chore: Remove unused duplicate main.py from shared pipeline * fix: Correct crosswalk paths in aarete_derived.py * chore: Remove unused documentation files from fieldExtraction * docs: Add documentation files to documentation folder * docs: Update README with uv setup, expanded project structure, and branching conventions * docs: Add uv installation steps with Ubuntu/WSL emphasis * Enable prompt caching for all remaining LLM calls - Add _INSTRUCTION() functions for: EXHIBIT_HEADER, EXHIBIT_LINKAGE, EXHIBIT_TITLE_MATCH, DATE_FIX, DERIVED_TERM_DATE, CHECK_PROVIDER_NAME_MATCH, SPECIAL_CASE_ASSIGNMENT - Update all invoke_claude() calls in saas and clover pipelines to use cache=True with corresponding _INSTRUCTION() functions - Add new instructions to get_cacheable_instructions() for cache warming - Update tests for new instruction functions Functions now using caching: - prompt_exhibit_level - prompt_exhibit_lesser (EXHIBIT_LEVEL_LESSER_OF) - prompt_fee_schedule_breakout - prompt_grouper_breakout - prompt_special_case_assignment - prompt_exhibit_linkage - prompt_exhibit_header - prompt_smart_chunked (ONE_TO_ONE templates) - prompt_date_fix - prompt_derived_term_date - prompt_exhibit_title_match - provider_name_match_check 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com> * Reorder * feat: Add bcbs_promise client pipeline with OFFSET_TERM extraction - Add new bcbs_promise client with HSC-based OFFSET_TERM field extraction - Extract full paragraph text of offset/recoupment provisions from contracts - Derive OFFSET_INDICATOR (Y/N) from OFFSET_TERM presence - Fix reorder_columns to preserve extra columns not in COLUMN_ORDER - Update QC/QA output path to outputs/qc_qa/ * fix: Update dev deps and test assertions for QC/QA output path - Add pytest/pytest-mock to dev dependencies for mypy type checking - Update test assertions to expect outputs/qc_qa instead of qa_qc_output * style: Apply black formatting to prompt_templates.py * Merge main, move scripts * Archive some scripts * update py version * remove .py version file * Remove ASCII characters * Restore testbed code * restore tracking * Update testbed metrics * Enable prompt caching for CODE_LAST_CHECK, FILL_BILL_TYPE, DUAL_LOB_CHECK, and GROUPER_BREAKOUT - Add CODE_LAST_CHECK_INSTRUCTION() for service specificity classification - Add FILL_BILL_TYPE_INSTRUCTION() for bill type code determination - Add DUAL_LOB_CHECK_INSTRUCTION() for Medicare/Medicaid classification - Update code_funcs.py to use caching for CODE_LAST_CHECK, FILL_BILL_TYPE, GROUPER_BREAKOUT - Update postprocessing_funcs.py to use caching for DUAL_LOB_CHECK - Add new instructions to get_cacheable_instructions() for cache warming - Add unit tests for new instruction functions 🤖 Generated with [Claude Code](https://claude.com/claude-code) Co-Authored-By: Claude Opus 4.5 <noreply@anthropic.com> * Fix postprocessing_funcs to remove invalid columns * Merge branch 'main' into feature/lesser-table-caching-refactor-hybrid * Revert prompt caching changes from aed1b73c * update formatting * Update imports Approved-by: Sha Brown Approved-by: Praneel Panchigar
537 lines
18 KiB
Python
537 lines
18 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,
|
|
USER_LIST,
|
|
)
|
|
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
|
|
|
|
REDIRECT_URI = "http://172.29.20.126:8503"
|
|
user_list = USER_LIST
|
|
|
|
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"].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?"])
|
|
)
|
|
|
|
field_row = st.columns([0.15, 0.45, 0.4])
|
|
with field_row[0]:
|
|
st.write("**Field Name**")
|
|
with field_row[1]:
|
|
field = st.selectbox(
|
|
"Field Name",
|
|
sorted(set(field_prompt_mapping.keys())),
|
|
index=0,
|
|
label_visibility="collapsed",
|
|
)
|
|
|
|
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]:
|
|
contract_count = st.selectbox(
|
|
"Contract count",
|
|
("1", "10", "20", "30", "50", "All"),
|
|
index=1,
|
|
label_visibility="collapsed",
|
|
)
|
|
|
|
seed_row = st.columns([0.15, 0.45, 0.4])
|
|
|
|
contract_list = sorted(os.listdir(SOURCE_DIRECTORY))
|
|
|
|
# to be deleted later
|
|
contract_list = [
|
|
contract
|
|
for contract in contract_list
|
|
if contract.replace(" MU", "").replace("_MU", "").replace(".txt", "")
|
|
in list(field_values["(internal) Document Name"])
|
|
]
|
|
|
|
with seed_row[0]:
|
|
if contract_count in ["10", "20", "30", "50"]:
|
|
st.write("**Seed Value**")
|
|
elif contract_count == "1":
|
|
st.write("**Contract Name**")
|
|
with seed_row[1]:
|
|
if contract_count in ["10", "20", "30", "50"]:
|
|
seed_value = st.text_input(
|
|
"**Seed Value**", value=20, label_visibility="collapsed"
|
|
)
|
|
random.seed(seed_value)
|
|
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"
|
|
)
|
|
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]:
|
|
llm_selected = st.selectbox(
|
|
"Langauge Model",
|
|
(
|
|
"Claude 2",
|
|
"Claude Instant",
|
|
"Llama 2 Chat 13B",
|
|
"Llama 2 Chat 70B",
|
|
"Titan Text Express",
|
|
),
|
|
label_visibility="collapsed",
|
|
)
|
|
|
|
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"):
|
|
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"
|
|
)
|
|
|
|
page_list_all = []
|
|
for contract in contract_list:
|
|
page_list = []
|
|
with open(
|
|
os.path.join(SOURCE_DIRECTORY, contract[:-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"{contract[:-4]}_page{page}.txt"
|
|
dict_with_pages = {"source": {"$eq": file_path}}
|
|
page_list.append(dict_with_pages)
|
|
page_list_all.append(page_list)
|
|
contract_txt_mapping = dict(zip(contract_list, page_list_all))
|
|
|
|
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")
|
|
field_values = field_values[column_list]
|
|
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")
|
|
|
|
# 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:
|
|
if llm_selected == "Titan Text Express":
|
|
LLM = Bedrock(
|
|
model_id="amazon.titan-text-express-v1",
|
|
client=bedrock_runtime,
|
|
model_kwargs={
|
|
"maxTokenCount": 512,
|
|
"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,
|
|
},
|
|
)
|
|
elif llm_selected == "Llama 2 Chat 13B":
|
|
LLM = Bedrock(
|
|
model_id="meta.llama2-13b-chat-v1",
|
|
client=bedrock_runtime,
|
|
model_kwargs={
|
|
"max_gen_len": 512,
|
|
"temperature": 0,
|
|
# "topP": 0.9,
|
|
},
|
|
)
|
|
elif llm_selected == "Claude Instant":
|
|
LLM = Bedrock(
|
|
model_id="anthropic.claude-instant-v1",
|
|
client=bedrock_runtime,
|
|
model_kwargs={
|
|
# "max_tokens_to_sample": 512,
|
|
"temperature": 0,
|
|
# "topP": 0.9,
|
|
},
|
|
)
|
|
elif llm_selected == "Claude 2":
|
|
LLM = Bedrock(
|
|
model_id="anthropic.claude-v2:1",
|
|
client=bedrock_runtime,
|
|
model_kwargs={
|
|
# "max_tokens_to_sample": 512,
|
|
"temperature": 0,
|
|
# "topP": 0.9,
|
|
},
|
|
)
|
|
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'])
|
|
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",
|
|
]
|
|
)
|
|
try:
|
|
history = pd.read_csv("history.csv")
|
|
except:
|
|
history = pd.DataFrame(
|
|
columns=[
|
|
"Field Name",
|
|
"# Contracts Tested",
|
|
"Username",
|
|
"Date/Time",
|
|
"Accuracy",
|
|
"Attempt #",
|
|
]
|
|
)
|
|
attempt = 0
|
|
|
|
if llm_selected in ["Llama 2 Chat 13B", "Llama 2 Chat 70B"]:
|
|
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:
|
|
RETRIEVER = st.session_state.DB.as_retriever(
|
|
search_kwargs={"filter": {"$or": page_list}, "k": k_value}
|
|
)
|
|
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},
|
|
)
|
|
score = st.session_state.DB.similarity_search_with_relevance_scores(
|
|
prompt, k=4, filter={"$or": page_list}
|
|
)
|
|
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)
|
|
|
|
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
|
|
elif 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 "Unfortunately, I do not have enough context" 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]
|
|
|
|
df["Contract Name"] = contract_list
|
|
# to be deleted later
|
|
df["Contract Name"] = [
|
|
contract.replace(" MU", "").replace("_MU", "").replace(".txt", "")
|
|
for contract in contract_list
|
|
]
|
|
|
|
df["New Extracted value"] = answer_list
|
|
df["Confidence Level"] = [round(score, 2) for score in score_list]
|
|
Snippet = []
|
|
count = 0
|
|
for i in range(int(contract_count)):
|
|
for j in range(10):
|
|
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]
|
|
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
|
|
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()]
|
|
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",
|
|
]
|
|
]
|
|
else:
|
|
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",
|
|
]
|
|
]
|
|
|
|
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,
|
|
]
|
|
# df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False)
|
|
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()
|
|
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")
|
|
|
|
st.write(column_name)
|
|
st.write(len(contract_list))
|
|
|
|
else:
|
|
st.write("Access Denied")
|