Merged DEV into feature/streamlit_update_connections
This commit is contained in:
@@ -194,7 +194,7 @@ resource "aws_instance" "streamlit_server" {
|
||||
Type=simple
|
||||
Restart=always
|
||||
WorkingDirectory=/home/ubuntu/doczy.ai/streamlit/multipage
|
||||
ExecStart=/home/ubuntu/.local/bin/streamlit run /home/ubuntu/doczy.ai/streamlit/multipage/full_interface.py --server.port 8505
|
||||
ExecStart=/home/ubuntu/.local/bin/streamlit run /home/ubuntu/doczy.ai/streamlit/multipage/Interface_0.py --server.port 8505
|
||||
|
||||
[Install]
|
||||
WantedBy=multi-user.target
|
||||
|
||||
@@ -0,0 +1,241 @@
|
||||
import json
|
||||
import security
|
||||
import streamlit as st
|
||||
from streamlit_extras.add_vertical_space import add_vertical_space
|
||||
import os
|
||||
import streamlit as st
|
||||
import pandas as pd
|
||||
from io import StringIO
|
||||
from datetime import datetime
|
||||
import boto3
|
||||
import util
|
||||
import requests
|
||||
from sf_conn import get_secret, save_to_sf
|
||||
from io import StringIO, BytesIO
|
||||
import time
|
||||
from sf_conn import get_client_names, insert_upload_logs
|
||||
from constants import USER_LIST
|
||||
|
||||
(REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(0)
|
||||
|
||||
user_list = USER_LIST
|
||||
if 'uploading' not in st.session_state:
|
||||
st.session_state.uploading = False
|
||||
if 'upload_key' not in st.session_state:
|
||||
st.session_state.upload_key = 0
|
||||
if 'file_list' not in st.session_state:
|
||||
st.session_state.file_list = []
|
||||
if 'show_batchID' not in st.session_state:
|
||||
st.session_state.show_batchID = False
|
||||
if 'landing_zone' not in st.session_state:
|
||||
st.session_state.landing_zone = 'contracts-landing-zone'
|
||||
if 'batch_id' not in st.session_state:
|
||||
st.session_state.batch_id = 'failed_cases'
|
||||
if 'client_bucket' not in st.session_state:
|
||||
st.session_state.client_bucket = 'default_bucket'
|
||||
|
||||
def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload_user):
|
||||
"""
|
||||
Input: batch_id, client_name, file_name, upload_datetime, upload_user
|
||||
Output: status of the insert query
|
||||
"""
|
||||
try:
|
||||
return 'Log inserted successfully'
|
||||
except Exception as e:
|
||||
return e
|
||||
|
||||
|
||||
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")
|
||||
|
||||
# AARETE LOGO
|
||||
x,y,z = st.columns([15,2,15])
|
||||
with y:
|
||||
st.image('aaretelogo.png')
|
||||
|
||||
hide_img_fs = '''
|
||||
<style>
|
||||
button[title="View fullscreen"]{
|
||||
visibility: hidden;}
|
||||
</style>
|
||||
'''
|
||||
st.markdown(hide_img_fs, unsafe_allow_html=True)
|
||||
|
||||
_,c1= st.columns([5,1])
|
||||
try:
|
||||
util.setup_page(REDIRECT_URI)
|
||||
except Exception as e:
|
||||
st.write(f"SSO Failed = {e}")
|
||||
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
|
||||
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.")
|
||||
auth_url = security.get_auth_url(REDIRECT_URI)
|
||||
st.markdown(f"<a href='{auth_url}' target='_self'>Sign In</a>", unsafe_allow_html=True)
|
||||
st.stop()
|
||||
|
||||
s3_client = boto3.client('s3',
|
||||
region_name="us-east-2",
|
||||
)
|
||||
|
||||
client_list = ['doczy-ai-client-1', 'Delaware First Health, Inc.', 'Community Health Choice, Inc','CareSource Network Partners LLC',
|
||||
'HealthNet of Cali', 'Oklahoma Complete Health, Inc', 'HealthFirst', 'Molina Healthcare of TX', 'AvMed', 'Arizona Care1st',
|
||||
'WellCare New Jersey']
|
||||
|
||||
# This is the list of client fetched from Snowflake
|
||||
# TODO: Need to update the streamlit code to use the client names from this list
|
||||
# And use the s3 paths to save the objects for the respective client
|
||||
client_list, s3_paths = get_client_names()
|
||||
client_s3_paths = dict(zip(client_list, s3_paths))
|
||||
|
||||
client_row = st.columns([0.1, 0.8])
|
||||
with client_row[0]:
|
||||
st.write("**Client Name**")
|
||||
with client_row[1]:
|
||||
client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed", index = None)
|
||||
|
||||
client_bucket = client_s3_paths.get(client)
|
||||
|
||||
# to be deleted when client buckets are created
|
||||
client_bucket = 'doczy-dev-infra-textract'
|
||||
|
||||
file_row = st.columns([0.1, 0.8])
|
||||
with file_row[0]:
|
||||
st.write("**Upload Files**")
|
||||
with file_row[1]:
|
||||
file_list = st.file_uploader("Upload", type=['docx','tiff','pdf'], accept_multiple_files=True, label_visibility = "collapsed", help="Only PDF, TIFF and DOCX file formats are supported.", disabled=st.session_state.uploading, key = st.session_state.upload_key)
|
||||
|
||||
add_vertical_space(2)
|
||||
df = pd.DataFrame(columns=['Contract Name'])
|
||||
df['Contract Name'] = file_list
|
||||
file_names = []
|
||||
buttons = st.columns([0.4, 0.4, 0.2])
|
||||
|
||||
def set_uploading_state():
|
||||
if not client == None and not len(file_list) == 0:
|
||||
st.session_state.file_list = file_list
|
||||
st.session_state.upload_key += 1
|
||||
st.session_state.uploading = True
|
||||
|
||||
with buttons[1]:
|
||||
if st.button("Create Batch", on_click=set_uploading_state):
|
||||
file_list = st.session_state.file_list
|
||||
if client == None:
|
||||
st.error("No Client Name Selected.")
|
||||
elif len(file_list) == 0:
|
||||
st.error("No Files Selected.")
|
||||
else:
|
||||
with st.spinner('Running...'):
|
||||
myobj = { "client-bucket-name": client_bucket }
|
||||
response = requests.post(create_batch_url, json = myobj)
|
||||
if response.status_code >= 200 and response.status_code < 300:
|
||||
try:
|
||||
batch_id = json.loads(json.loads(response.text)['body'])['batch_id']
|
||||
landing_zone = json.loads(json.loads(response.text)['body'])['landing_zone']
|
||||
except:
|
||||
# st.write(myobj)
|
||||
# st.write(response.text)
|
||||
st.write("Internal Error. Reach out to Doczy.AI Team.")
|
||||
batch_id = 'failed_cases'
|
||||
landing_zone = 'contracts-landing-zone'
|
||||
else:
|
||||
st.error("Failed")
|
||||
for uploaded_file in file_list:
|
||||
stringio = BytesIO(uploaded_file.getvalue())
|
||||
stringio.seek(0)
|
||||
s3_client.put_object(Bucket=client_bucket, Body=stringio.getvalue(), Key=
|
||||
landing_zone+batch_id+'/'+str(uploaded_file.name))
|
||||
# s3_client.put_object_tagging(Bucket=client_bucket, Key=
|
||||
# landing_zone+batch_id+'/'+str(uploaded_file.name), Tagging = {'TagSet': [ { 'Key': 'BatchId', 'Value': batch_id }]})
|
||||
|
||||
# TODO: Test this insert function with snowflake
|
||||
upload_log = insert_upload_logs(batch_id, client, str(uploaded_file.name), datetime.now().strftime("%Y-%m-%d %H:%M:%S"), user_mail)
|
||||
st.write(upload_log)
|
||||
|
||||
file_names.append(str(uploaded_file.name))
|
||||
st.session_state.uploading = False
|
||||
st.session_state.show_batchID = True
|
||||
st.session_state.client_bucket = client_bucket
|
||||
st.session_state.batch_id = batch_id
|
||||
st.session_state.landing_zone = landing_zone
|
||||
st.rerun()
|
||||
if st.session_state.show_batchID:
|
||||
batch_id = st.session_state.batch_id
|
||||
client_bucket = st.session_state.client_bucket
|
||||
landing_zone = st.session_state.landing_zone
|
||||
st.write("Batch ID Created - ")
|
||||
st.code(f"{batch_id}")
|
||||
st.write(f"Files uploaded to s3://{client_bucket}/{landing_zone}{batch_id}")
|
||||
st.session_state.show_batchID = False
|
||||
st.session_state.file_list = []
|
||||
st.session_state.batch_id = 'failed_cases'
|
||||
st.session_state.client_bucket = 'default_bucket'
|
||||
st.session_state.landing_zone = 'contracts-landing-zone'
|
||||
|
||||
|
||||
|
||||
# @st.cache_data
|
||||
# def convert_df(df):
|
||||
# return df.to_csv(index=False).encode('utf-8')
|
||||
|
||||
# csv = convert_df(df)
|
||||
|
||||
# df = df.reset_index() # make sure indexes pair with number of rows
|
||||
|
||||
# contract_list = []
|
||||
# for index, row in df.iterrows():
|
||||
# group_list = []
|
||||
# if row['Unique Key']:
|
||||
# group_list.append('Unique Key')
|
||||
# if row['Pricing Before Carveouts']:
|
||||
# group_list.append('Pricing Before Carveouts')
|
||||
# if row['Contract Related']:
|
||||
# group_list.append('Contract Related')
|
||||
# if row['Provider']:
|
||||
# group_list.append('Provider')
|
||||
# if row['Timeline']:
|
||||
# group_list.append('Timeline')
|
||||
# if row['Carveout Indicator']:
|
||||
# group_list.append('Carveout Indicator')
|
||||
# if row['Carveout Methodology']:
|
||||
# group_list.append('Carveout Methodology')
|
||||
# entry_dict = {
|
||||
# "contract_name": row['Contract Name'],
|
||||
# "groups": group_list,
|
||||
# "contract_source_path": "batches/batch_1/"+client+"/"+row['Contract Name']
|
||||
# }
|
||||
# contract_list.append(entry_dict)
|
||||
# contract_list = list(df['Contract Name'])
|
||||
|
||||
# myobj = {
|
||||
# "s3_bucket": 'doczy-dev-infra-textract',
|
||||
# "batch_id": "1",
|
||||
# "client_name": client,
|
||||
# "username": user_mail,
|
||||
# "contract_list": contract_list
|
||||
# }
|
||||
|
||||
# buttons = st.columns([0.8, 0.2])
|
||||
# with buttons[0]:
|
||||
# st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
|
||||
# with buttons[1]:
|
||||
# if st.button("Upload to DB"):
|
||||
# response = requests.post(doczy_pipeline, json = myobj)
|
||||
# if response.status_code >= 200 and response.status_code < 300:
|
||||
# st.write("Success")
|
||||
# else:
|
||||
# st.write("Failed")
|
||||
# # st.write(response.text)
|
||||
Binary file not shown.
|
After Width: | Height: | Size: 68 KiB |
@@ -0,0 +1,228 @@
|
||||
import os
|
||||
|
||||
# from dotenv import load_dotenv
|
||||
from chromadb.config import Settings
|
||||
|
||||
# https://python.langchain.com/en/latest/modules/indexes/document_loaders/examples/excel.html?highlight=xlsx#microsoft-excel
|
||||
from langchain_community.document_loaders import CSVLoader, PDFMinerLoader, TextLoader, UnstructuredExcelLoader, Docx2txtLoader
|
||||
from langchain_community.document_loaders import UnstructuredFileLoader, UnstructuredMarkdownLoader
|
||||
|
||||
|
||||
|
||||
# load_dotenv()
|
||||
# ROOT_DIRECTORY = os.path.dirname(os.path.realpath(__file__))
|
||||
ROOT_DIRECTORY = "\\\\amznfsxuofkyi1z.aarete.local\\SharedFiles\\AArete Client Work\\Modahealth\\Restricted\\Moda Growth\\Artificial Intelligence\\DEFAXXER_20231207"
|
||||
|
||||
# Define the folder for source and output
|
||||
SOURCE_DIRECTORY = "SOURCE_DOCUMENTS"
|
||||
OUTPUT_DIRECTORY = f"{ROOT_DIRECTORY}\\Output"
|
||||
|
||||
PERSIST_DIRECTORY = 'DB'
|
||||
|
||||
MODELS_PATH = "C:\\Users\\Public\\models"
|
||||
|
||||
# Can be changed to a specific number
|
||||
INGEST_THREADS = os.cpu_count() or 8
|
||||
|
||||
# Define the Chroma settings
|
||||
CHROMA_SETTINGS = Settings(
|
||||
anonymized_telemetry=False,
|
||||
is_persistent=True,
|
||||
allow_reset=True,
|
||||
)
|
||||
|
||||
# Context Window and Max New Tokens
|
||||
CONTEXT_WINDOW_SIZE = 4096
|
||||
MAX_NEW_TOKENS = CONTEXT_WINDOW_SIZE # int(CONTEXT_WINDOW_SIZE/4)
|
||||
|
||||
#### If you get a "not enough space in the buffer" error, you should reduce the values below, start with half of the original values and keep halving the value until the error stops appearing
|
||||
|
||||
N_GPU_LAYERS = 100 # Llama-2-70B has 83 layers
|
||||
N_BATCH = 512
|
||||
|
||||
### From experimenting with the Llama-2-7B-Chat-GGML model on 8GB VRAM, these values work:
|
||||
# N_GPU_LAYERS = 20
|
||||
# N_BATCH = 512
|
||||
|
||||
|
||||
# https://python.langchain.com/en/latest/_modules/langchain/document_loaders/excel.html#UnstructuredExcelLoader
|
||||
DOCUMENT_MAP = {
|
||||
".txt": TextLoader,
|
||||
".md": UnstructuredMarkdownLoader,
|
||||
".py": TextLoader,
|
||||
# ".pdf": PDFMinerLoader,
|
||||
".pdf": UnstructuredFileLoader,
|
||||
".csv": CSVLoader,
|
||||
".xls": UnstructuredExcelLoader,
|
||||
".xlsx": UnstructuredExcelLoader,
|
||||
".docx": Docx2txtLoader,
|
||||
".doc": Docx2txtLoader,
|
||||
}
|
||||
|
||||
# Default Instructor Model
|
||||
EMBEDDING_MODEL_NAME = "hkunlp/instructor-large" # Uses 1.5 GB of VRAM (High Accuracy with lower VRAM usage)
|
||||
|
||||
####
|
||||
#### OTHER EMBEDDING MODEL OPTIONS
|
||||
####
|
||||
|
||||
# EMBEDDING_MODEL_NAME = "hkunlp/instructor-xl" # Uses 5 GB of VRAM (Most Accurate of all models)
|
||||
# EMBEDDING_MODEL_NAME = "intfloat/e5-large-v2" # Uses 1.5 GB of VRAM (A little less accurate than instructor-large)
|
||||
# EMBEDDING_MODEL_NAME = "intfloat/e5-base-v2" # Uses 0.5 GB of VRAM (A good model for lower VRAM GPUs)
|
||||
# EMBEDDING_MODEL_NAME = "all-MiniLM-L6-v2" # Uses 0.2 GB of VRAM (Less accurate but fastest - only requires 150mb of vram)
|
||||
|
||||
####
|
||||
#### MULTILINGUAL EMBEDDING MODELS
|
||||
####
|
||||
|
||||
# EMBEDDING_MODEL_NAME = "intfloat/multilingual-e5-large" # Uses 2.5 GB of VRAM
|
||||
# EMBEDDING_MODEL_NAME = "intfloat/multilingual-e5-base" # Uses 1.2 GB of VRAM
|
||||
|
||||
|
||||
#### SELECT AN OPEN SOURCE LLM (LARGE LANGUAGE MODEL)
|
||||
# Select the Model ID and model_basename
|
||||
# load the LLM for generating Natural Language responses
|
||||
|
||||
#### GPU VRAM Memory required for LLM Models (ONLY) by Billion Parameter value (B Model)
|
||||
#### Does not include VRAM used by Embedding Models - which use an additional 2GB-7GB of VRAM depending on the model.
|
||||
####
|
||||
#### (B Model) (float32) (float16) (GPTQ 8bit) (GPTQ 4bit)
|
||||
#### 7b 28 GB 14 GB 7 GB - 9 GB 3.5 GB - 5 GB
|
||||
#### 13b 52 GB 26 GB 13 GB - 15 GB 6.5 GB - 8 GB
|
||||
#### 32b 130 GB 65 GB 32.5 GB - 35 GB 16.25 GB - 19 GB
|
||||
#### 65b 260.8 GB 130.4 GB 65.2 GB - 67 GB 32.6 GB - - 35 GB
|
||||
|
||||
# MODEL_ID = "TheBloke/Llama-2-7B-Chat-GGML"
|
||||
# MODEL_BASENAME = "llama-2-7b-chat.ggmlv3.q4_0.bin"
|
||||
|
||||
####
|
||||
#### (FOR GGUF MODELS)
|
||||
####
|
||||
|
||||
# MODEL_ID = "TheBloke/Llama-2-13b-Chat-GGUF"
|
||||
# MODEL_BASENAME = "llama-2-13b-chat.Q4_K_M.gguf"
|
||||
|
||||
MODEL_ID = "TheBloke/Llama-2-7b-Chat-GGUF"
|
||||
MODEL_BASENAME = "llama-2-7b-chat.Q4_K_M.gguf"
|
||||
|
||||
# MODEL_ID = "TheBloke/Mistral-7B-Instruct-v0.1-GGUF"
|
||||
# MODEL_BASENAME = "mistral-7b-instruct-v0.1.Q8_0.gguf"
|
||||
|
||||
# MODEL_ID = "TheBloke/Llama-2-70b-Chat-GGUF"
|
||||
# MODEL_BASENAME = "llama-2-70b-chat.Q4_K_M.gguf"
|
||||
|
||||
####
|
||||
#### (FOR HF MODELS)
|
||||
####
|
||||
|
||||
# MODEL_ID = "NousResearch/Llama-2-7b-chat-hf"
|
||||
# MODEL_BASENAME = None
|
||||
# MODEL_ID = "TheBloke/vicuna-7B-1.1-HF"
|
||||
# MODEL_BASENAME = None
|
||||
# MODEL_ID = "TheBloke/Wizard-Vicuna-7B-Uncensored-HF"
|
||||
# MODEL_ID = "TheBloke/guanaco-7B-HF"
|
||||
# MODEL_ID = 'NousResearch/Nous-Hermes-13b' # Requires ~ 23GB VRAM. Using STransformers
|
||||
# alongside will 100% create OOM on 24GB cards.
|
||||
# llm = load_model(device_type, model_id=model_id)
|
||||
|
||||
####
|
||||
#### (FOR GPTQ QUANTIZED) Select a llm model based on your GPU and VRAM GB. Does not include Embedding Models VRAM usage.
|
||||
####
|
||||
|
||||
##### 48GB VRAM Graphics Cards (RTX 6000, RTX A6000 and other 48GB VRAM GPUs) #####
|
||||
|
||||
### 65b GPTQ LLM Models for 48GB GPUs (*** With best embedding model: hkunlp/instructor-xl ***)
|
||||
# MODEL_ID = "TheBloke/guanaco-65B-GPTQ"
|
||||
# MODEL_BASENAME = "model.safetensors"
|
||||
# MODEL_ID = "TheBloke/Airoboros-65B-GPT4-2.0-GPTQ"
|
||||
# MODEL_BASENAME = "model.safetensors"
|
||||
# MODEL_ID = "TheBloke/gpt4-alpaca-lora_mlp-65B-GPTQ"
|
||||
# MODEL_BASENAME = "model.safetensors"
|
||||
# MODEL_ID = "TheBloke/Upstage-Llama1-65B-Instruct-GPTQ"
|
||||
# MODEL_BASENAME = "model.safetensors"
|
||||
|
||||
##### 24GB VRAM Graphics Cards (RTX 3090 - RTX 4090 (35% Faster) - RTX A5000 - RTX A5500) #####
|
||||
|
||||
### 13b GPTQ Models for 24GB GPUs (*** With best embedding model: hkunlp/instructor-xl ***)
|
||||
# MODEL_ID = "TheBloke/Wizard-Vicuna-13B-Uncensored-GPTQ"
|
||||
# MODEL_BASENAME = "Wizard-Vicuna-13B-Uncensored-GPTQ-4bit-128g.compat.no-act-order.safetensors"
|
||||
# MODEL_ID = "TheBloke/vicuna-13B-v1.5-GPTQ"
|
||||
# MODEL_BASENAME = "model.safetensors"
|
||||
# MODEL_ID = "TheBloke/Nous-Hermes-13B-GPTQ"
|
||||
# MODEL_BASENAME = "nous-hermes-13b-GPTQ-4bit-128g.no-act.order"
|
||||
# MODEL_ID = "TheBloke/WizardLM-13B-V1.2-GPTQ"
|
||||
# MODEL_BASENAME = "gptq_model-4bit-128g.safetensors
|
||||
|
||||
### 30b GPTQ Models for 24GB GPUs (*** Requires using intfloat/e5-base-v2 instead of hkunlp/instructor-large as embedding model ***)
|
||||
# MODEL_ID = "TheBloke/Wizard-Vicuna-30B-Uncensored-GPTQ"
|
||||
# MODEL_BASENAME = "Wizard-Vicuna-30B-Uncensored-GPTQ-4bit--1g.act.order.safetensors"
|
||||
# MODEL_ID = "TheBloke/WizardLM-30B-Uncensored-GPTQ"
|
||||
# MODEL_BASENAME = "WizardLM-30B-Uncensored-GPTQ-4bit.act-order.safetensors"
|
||||
|
||||
##### 8-10GB VRAM Graphics Cards (RTX 3080 - RTX 3080 Ti - RTX 3070 Ti - 3060 Ti - RTX 2000 Series, Quadro RTX 4000, 5000, 6000) #####
|
||||
### (*** Requires using intfloat/e5-small-v2 instead of hkunlp/instructor-large as embedding model ***)
|
||||
|
||||
### 7b GPTQ Models for 8GB GPUs
|
||||
# MODEL_ID = "TheBloke/Wizard-Vicuna-7B-Uncensored-GPTQ"
|
||||
# MODEL_BASENAME = "Wizard-Vicuna-7B-Uncensored-GPTQ-4bit-128g.no-act.order.safetensors"
|
||||
# MODEL_ID = "TheBloke/WizardLM-7B-uncensored-GPTQ"
|
||||
# MODEL_BASENAME = "WizardLM-7B-uncensored-GPTQ-4bit-128g.compat.no-act-order.safetensors"
|
||||
# MODEL_ID = "TheBloke/wizardLM-7B-GPTQ"
|
||||
# MODEL_BASENAME = "wizardLM-7B-GPTQ-4bit.compat.no-act-order.safetensors"
|
||||
|
||||
####
|
||||
#### (FOR GGML) (Quantized cpu+gpu+mps) models - check if they support llama.cpp
|
||||
####
|
||||
|
||||
# MODEL_ID = "TheBloke/wizard-vicuna-13B-GGML"
|
||||
# MODEL_BASENAME = "wizard-vicuna-13B.ggmlv3.q4_0.bin"
|
||||
# MODEL_BASENAME = "wizard-vicuna-13B.ggmlv3.q6_K.bin"
|
||||
# MODEL_BASENAME = "wizard-vicuna-13B.ggmlv3.q2_K.bin"
|
||||
# MODEL_ID = "TheBloke/orca_mini_3B-GGML"
|
||||
# MODEL_BASENAME = "orca-mini-3b.ggmlv3.q4_0.bin"
|
||||
|
||||
####
|
||||
#### (FOR AWQ QUANTIZED) Select a llm model based on your GPU and VRAM GB. Does not include Embedding Models VRAM usage.
|
||||
### (*** MODEL_BASENAME is not actually used but have to contain .awq so the correct model loading is used ***)
|
||||
### (*** Compute capability 7.5 (sm75) and CUDA Toolkit 11.8+ are required ***)
|
||||
####
|
||||
# MODEL_ID = "TheBloke/Llama-2-7B-Chat-AWQ"
|
||||
# MODEL_BASENAME = "model.safetensors.awq"
|
||||
|
||||
|
||||
|
||||
##########################################################################################################################################
|
||||
|
||||
## CONSTANTS FOR INFRATRUCTURE
|
||||
|
||||
# SSO User list
|
||||
USER_LIST = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com', 'pkatariya@aarete.com'
|
||||
, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com'
|
||||
, 'vnair@aarete.com', 'kminhas@aarete.com', 'fmohiuddin@aarete.com', 'slitewka@aarete.com', 'qdoest@aarete.com', 'bkoryga@aarete.com', 'bcielecki@aarete.com', 'mszymanski@aarete.com','hupreti@aarete.com',
|
||||
'sshingare@aarete.com', 'vsrinivasan@aarete.com' ]
|
||||
|
||||
|
||||
# DOCZY DEV
|
||||
DOCZY_PIPELINE_URL_DEV = 'https://8ir4vi1ri4.execute-api.us-east-2.amazonaws.com/dev/'
|
||||
DOCZY_REDIRECT_URL_DEV = 'https://doczydev.aarete.com:850'
|
||||
DOCZY_CREATE_BATCH_URL_DEV = 'https://lfksus2t62.execute-api.us-east-2.amazonaws.com/dev/create-batch'
|
||||
|
||||
# DOCZY UAT
|
||||
DOCZY_PIPELINE_URL_UAT = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/trigger-pipeline'
|
||||
DOCZY_REDIRECT_URL_UAT = 'https://doczyuat.aarete.com:850'
|
||||
DOCZY_CREATE_BATCH_URL_UAT = 'https://29gm8cek03.execute-api.us-east-2.amazonaws.com/dev/create-batch'
|
||||
|
||||
# DOCZY UAT
|
||||
|
||||
# SNOWFLAKE DEV DATABASE
|
||||
SNOWFLAKE_ACCOUNT_LOCATOR="aarete-doczyai",
|
||||
DEV_DB_ROLE = "DEVADMIN",
|
||||
DEV_WH="DEV_XS",
|
||||
DEV_DB="DOCZY_DEV",
|
||||
DEV_STAGING_SCHEMA="STG"
|
||||
|
||||
# SNOWFLAKE UAT DATABASE
|
||||
UAT_DB_ROLE = "UATADMIN",
|
||||
UAT_WH="UAT_XS",
|
||||
UAT_DB="DOCZY_UAT",
|
||||
UAT_STAGING_SCHEMA="STG"
|
||||
@@ -0,0 +1,277 @@
|
||||
import streamlit as st
|
||||
import security
|
||||
from streamlit_extras.add_vertical_space import add_vertical_space
|
||||
import os
|
||||
import streamlit as st
|
||||
import pandas as pd
|
||||
from io import StringIO
|
||||
from datetime import datetime
|
||||
import boto3
|
||||
import util
|
||||
import requests
|
||||
import time
|
||||
from sf_conn import get_client_names, get_secret, save_to_sf
|
||||
from constants import USER_LIST
|
||||
|
||||
(REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(1)
|
||||
|
||||
REDIRECT_URI = 'https://doczydev.aarete.com:8501'
|
||||
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")
|
||||
|
||||
# AARETE LOGO
|
||||
x,y,z = st.columns([15,2,15])
|
||||
with y:
|
||||
st.image('aaretelogo.png')
|
||||
|
||||
hide_img_fs = '''
|
||||
<style>
|
||||
button[title="View fullscreen"]{
|
||||
visibility: hidden;}
|
||||
</style>
|
||||
'''
|
||||
st.markdown(hide_img_fs, unsafe_allow_html=True)
|
||||
|
||||
_,c1= st.columns([5,1])
|
||||
try:
|
||||
util.setup_page(REDIRECT_URI)
|
||||
except:
|
||||
st.write("SSO Failed")
|
||||
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
|
||||
try:
|
||||
c1.write(f"User: **{st.session_state.user_info['displayName']}**")
|
||||
user_mail = st.session_state.user_info['mail']
|
||||
except KeyError as e:
|
||||
# Do we add a link to get to the login page here?
|
||||
st.write("Session Expired.")
|
||||
# st.write("Please sign-in to use this app.")
|
||||
auth_url = security.get_auth_url(REDIRECT_URI)
|
||||
st.markdown(f"<a href='{auth_url}' target='_self'>Sign In</a>", unsafe_allow_html=True)
|
||||
st.stop()
|
||||
|
||||
|
||||
|
||||
s3_client = boto3.client('s3',
|
||||
region_name="us-east-2",
|
||||
)
|
||||
|
||||
# # to be replaced with snowflake data
|
||||
# client_list = ['doczy-ai-client-1', 'Delaware First Health, Inc.', 'Community Health Choice, Inc','CareSource Network Partners LLC',
|
||||
# 'HealthNet of Cali', 'Oklahoma Complete Health, Inc', 'HealthFirst', 'Molina Healthcare of TX', 'AvMed', 'Arizona Care1st',
|
||||
# 'WellCare New Jersey']
|
||||
|
||||
client_list, s3_paths = get_client_names()
|
||||
client_s3_paths = dict(zip(client_list, s3_paths))
|
||||
|
||||
client_row = st.columns([0.1, 0.8])
|
||||
with client_row[0]:
|
||||
st.write("**Client Name**")
|
||||
with client_row[1]:
|
||||
client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed", index = None) # MODIFIED for Ticket DOC-344
|
||||
|
||||
client_bucket = client_s3_paths.get(client)
|
||||
|
||||
# to be deleted when buckets for different clients are ready; below line is added only for testing the corresponding DAG
|
||||
client_bucket = 'doczy-dev-infra-textract'
|
||||
|
||||
batch_objects = s3_client.list_objects_v2(Bucket=client_bucket
|
||||
, Prefix="contracts-landing-zone/", Delimiter='/')
|
||||
|
||||
batch_list = []
|
||||
for prefix in batch_objects['CommonPrefixes']:
|
||||
batch_list.append(prefix['Prefix'][:-1].split('/')[-1])
|
||||
|
||||
# Hardcoded batch_list for testing purposes
|
||||
# batch_list = ['batch_020524103737', 'batch_090524131433', 'batch_090524131607', 'batch_100524123000', 'batch_130524064322',
|
||||
# 'batch_160524071331', 'batch_200524213550', 'batch_250424112237', 'batch_280524120530', 'batch_280524121721', 'batch_280524144222',
|
||||
# 'batch_290524123926', 'batch_290524164044', 'batch_310524102029', 'batch_310524124050', 'batch_310524162346', 'batch_310524162631']
|
||||
|
||||
if 'sorted_list' not in st.session_state:
|
||||
st.session_state.sorted_list = batch_list
|
||||
|
||||
def sort_list(ex_list, sort_by, order):
|
||||
if sort_by == 'Alphabetical':
|
||||
ex_list = sorted(ex_list, reverse=(order == 'Descending'))
|
||||
elif sort_by == 'Create Date':
|
||||
ex_list = ex_list if order == 'Ascending' else list(reversed(ex_list))
|
||||
|
||||
return ex_list
|
||||
|
||||
col1, col2, col3, col4 = st.columns([0.5, 0.5, 0.5, 0.5])
|
||||
|
||||
|
||||
with col1:
|
||||
sort_by = st.radio("**Sort Batch_IDs**", ('Alphabetical', 'Create Date'))
|
||||
|
||||
with col2:
|
||||
order = st.radio('', ('Ascending','Descending'))
|
||||
|
||||
with col3:
|
||||
add_vertical_space(2)
|
||||
if st.button('Apply'):
|
||||
st.session_state.sorted_list = sort_list(batch_list, sort_by, order)
|
||||
|
||||
path_row = st.columns([0.1, 0.8])
|
||||
with path_row[0]:
|
||||
st.write("**Batch ID**")
|
||||
with path_row[1]:
|
||||
batch_id = st.selectbox('**Batch ID**', st.session_state.sorted_list, label_visibility = "collapsed", index = None) # MODIFIED for Ticket DOC-344
|
||||
|
||||
if not batch_id:
|
||||
batch_id = "None"
|
||||
|
||||
checks = st.columns([0.1, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12, 0.12])
|
||||
with checks[0]:
|
||||
st.write("**Group No.**")
|
||||
|
||||
with checks[1]:
|
||||
a = st.checkbox('Unique Key', key = str(1), args="Unique")
|
||||
with checks[2]:
|
||||
b = st.checkbox('Pricing Before Carveouts', key = str(2))
|
||||
with checks[3]:
|
||||
c = st.checkbox('Contract Related', key = str(3))
|
||||
with checks[4]:
|
||||
d = st.checkbox('Provider', key = str(4))
|
||||
with checks[5]:
|
||||
e = st.checkbox('Timeline', key = str(5))
|
||||
with checks[6]:
|
||||
f = st.checkbox('Carveout Indicator', key = str(6))
|
||||
with checks[7]:
|
||||
g = st.checkbox('Carveout Methodology', key = str(7))
|
||||
|
||||
add_vertical_space(1)
|
||||
|
||||
df = pd.DataFrame(columns=['Contract Name', 'Unique Key','Pricing Before Carveouts'
|
||||
, 'Contract Related', 'Provider', 'Timeline', 'Carveout Indicator', 'Carveout Methodology'])
|
||||
file_list = []
|
||||
file_objects = s3_client.list_objects_v2(Bucket=client_bucket
|
||||
, Prefix="contracts_landing_zone/"+batch_id+"/", Delimiter='/')
|
||||
|
||||
# Hardcoded file_list for testing purposes
|
||||
# file_list = ['Boilerplate_TX Amendment Mission Health Network effective_040114 MU.pdf', 'Custom_TX - MP AMENDMENT - MISSION HEALTH NETWORK - MU.pdf',
|
||||
# 'Delaware First Health_First State Homecare Agency_212260_7 MU.pdf', 'Molina Healthcare of Texas, Inc. Amendment 4 - HIX ACA__EFF 01012016_MU.pdf']
|
||||
|
||||
if st.button("Read the contracts from Path"):
|
||||
for obj in file_objects.get('Contents',[]):
|
||||
if not obj['Key'].endswith('/'):
|
||||
file_list.append(obj['Key'].split('/')[-1])
|
||||
|
||||
df['Contract Name'] = file_list
|
||||
# df['Request ID'] = range(len(file_list))
|
||||
# df['Contract ID'] = file_list
|
||||
df['Unique Key'] = a
|
||||
df['Pricing Before Carveouts'] = b
|
||||
df['Contract Related'] = c
|
||||
df['Provider'] = d
|
||||
df['Timeline'] = e
|
||||
df['Carveout Indicator'] = f
|
||||
df['Carveout Methodology'] = g
|
||||
dir_path = os.path.dirname(os.path.realpath(__file__))
|
||||
print(f'DEBUGGING: PWD= {dir_path}')
|
||||
df.to_csv('temp1.csv', index=False)
|
||||
|
||||
add_vertical_space(1)
|
||||
df2 = pd.read_csv('temp1.csv')
|
||||
edited_df = st.data_editor(df2)
|
||||
|
||||
edited_df['REQUEST_USER'] = user_mail
|
||||
edited_df['LATEST_FLAG BOOLEAN'] = True
|
||||
edited_df['PIPELINE_KICKOFF_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
edited_df['REQUEST_DATETIME'] = datetime.now().strftime("%Y-%m-%d %H:%M:%S")
|
||||
|
||||
@st.cache_data
|
||||
def convert_df(df):
|
||||
return df.to_csv(index=False).encode('utf-8')
|
||||
|
||||
csv = convert_df(edited_df)
|
||||
# edited_df = edited_df.reset_index() # make sure indexes pair with number of rows
|
||||
|
||||
# additional_info = pd.DataFrame(columns=['REQUEST_ID','T_DRIVE_PATH','CLIENT_NAME'
|
||||
# , 'GROUP_NAME', 'REQUEST_USERNAME', 'REQUEST_DATETIME'])
|
||||
additional_info = pd.DataFrame(columns=['CLIENT_NAME', 'BATCH_ID', 'REQUEST_USERNAME', 'REQUEST_DATETIME'])
|
||||
additional_info.loc[0] = [client, batch_id, st.session_state.user_info['mail'], datetime.now().strftime("%Y-%m-%d %H:%M:%S")]
|
||||
st.write(additional_info)
|
||||
|
||||
st.session_state.contract_count = 0
|
||||
contract_list = []
|
||||
for index, row in edited_df.iterrows():
|
||||
allow_run_for_contract = False
|
||||
group_list = []
|
||||
if row['Unique Key']:
|
||||
group_list.append('A')
|
||||
allow_run_for_contract = True
|
||||
if row['Pricing Before Carveouts']:
|
||||
group_list.append('B')
|
||||
allow_run_for_contract = True
|
||||
if row['Contract Related']:
|
||||
group_list.append('C')
|
||||
allow_run_for_contract = True
|
||||
if row['Provider']:
|
||||
group_list.append('D')
|
||||
allow_run_for_contract = True
|
||||
if row['Timeline']:
|
||||
group_list.append('E')
|
||||
allow_run_for_contract = True
|
||||
if row['Carveout Indicator']:
|
||||
group_list.append('F')
|
||||
allow_run_for_contract = True
|
||||
if row['Carveout Methodology']:
|
||||
group_list.append('G')
|
||||
allow_run_for_contract = True
|
||||
if allow_run_for_contract: st.session_state.contract_count += 1
|
||||
entry_dict = {
|
||||
"contract_name": row['Contract Name'],
|
||||
"groups": group_list,
|
||||
"contract_source_path": "contracts_landing_zone/"+batch_id+"/"+row['Contract Name']
|
||||
}
|
||||
contract_list.append(entry_dict)
|
||||
|
||||
myobj = {
|
||||
"s3_bucket": client_bucket,
|
||||
"batch_id": batch_id,
|
||||
"client_name": client,
|
||||
"username": user_mail,
|
||||
"contract_list": contract_list
|
||||
}
|
||||
|
||||
buttons = st.columns([0.8, 0.2])
|
||||
with buttons[0]:
|
||||
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
|
||||
with buttons[1]:
|
||||
if st.button("Run Doczy.AI Pipeline"):
|
||||
if not st.session_state.contract_count == len(edited_df):
|
||||
st.error("Select at least one Group No. for every Contract")
|
||||
else:
|
||||
with st.spinner('Running...'):
|
||||
# csv_buf = StringIO()
|
||||
# additional_info.to_csv(csv_buf, header=True, index=False)
|
||||
# csv_buf.seek(0)
|
||||
# s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/request_submission.csv')
|
||||
# csv_buf = StringIO()
|
||||
# edited_df.to_csv(csv_buf, header=True, index=False)
|
||||
# csv_buf.seek(0)
|
||||
# s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/contract_config.csv')
|
||||
# try:
|
||||
# save_to_sf('load_request_and_contract_submissions', request_submission_file_name = "request_submission.csv", contract_config_file_name = "contract_config.csv")
|
||||
# except Exception as e:
|
||||
# st.write(e)
|
||||
response = requests.post(doczy_pipeline, json = myobj)
|
||||
if response.status_code >= 200 and response.status_code < 300:
|
||||
st.write("Success")
|
||||
# st.write(myobj)
|
||||
else:
|
||||
st.write("Failed")
|
||||
# st.write(response.text)
|
||||
|
||||
@@ -0,0 +1,267 @@
|
||||
import json
|
||||
import security
|
||||
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
|
||||
from streamlit_pdf_viewer import pdf_viewer # needs to be installed on the server
|
||||
import os
|
||||
import pandas as pd
|
||||
import numpy as np
|
||||
import util
|
||||
import anthropic
|
||||
from pydantic import BaseModel
|
||||
from typing import List
|
||||
import re
|
||||
import base64
|
||||
from sf_conn import get_snowflake_conn
|
||||
from sf_conn import get_client_names, get_secret, save_to_sf
|
||||
import io
|
||||
|
||||
(REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(2)
|
||||
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")
|
||||
|
||||
# AARETE LOGO
|
||||
x,y,z = st.columns([15,2,15])
|
||||
with y:
|
||||
st.image('aaretelogo.png')
|
||||
|
||||
hide_img_fs = '''
|
||||
<style>
|
||||
button[title="View fullscreen"]{
|
||||
visibility: hidden;}
|
||||
</style>
|
||||
'''
|
||||
st.markdown(hide_img_fs, unsafe_allow_html=True)
|
||||
|
||||
_,c1= st.columns([5,1])
|
||||
try:
|
||||
util.setup_page(REDIRECT_URI)
|
||||
except:
|
||||
st.write("SSO Failed")
|
||||
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
|
||||
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.write("Please sign-in to use this app.")
|
||||
auth_url = security.get_auth_url(REDIRECT_URI)
|
||||
st.markdown(f"<a href='{auth_url}' target='_self'>Sign In</a>", unsafe_allow_html=True)
|
||||
st.stop()
|
||||
|
||||
# remove below try except statement if comparison with actual vales is not required
|
||||
try:
|
||||
conn = get_snowflake_conn('STG')
|
||||
cur = conn.cursor()
|
||||
query = 'select * from "TRAINING_DATA_RAW"'
|
||||
cur.execute(query)
|
||||
field_values = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
|
||||
field_values['Document_Name'] = field_values['DOCUMENT_NAME']
|
||||
except:
|
||||
# field_values = pd.read_csv('contract_field_values.csv', encoding='unicode_escape', skipinitialspace=True)
|
||||
# field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:')]
|
||||
st.write("Conn failed, unable to fetch data from training data table in Snowflake")
|
||||
|
||||
try:
|
||||
query = 'select * from "PROMPT_CONFIG"'
|
||||
cur.execute(query)
|
||||
fields = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
|
||||
|
||||
fields.rename(columns={'FIELD_DESC': 'Field Name'}, inplace = True)
|
||||
fields.rename(columns={'PROMPT': 'Interrogation Question?'}, inplace = True)
|
||||
fields.rename(columns={'GROUP_ID': 'PRIORITY'}, inplace = True)
|
||||
fields.rename(columns={'FIELD_NAME': 'SF_DB_COL_NAME'}, inplace = True)
|
||||
fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True)
|
||||
except Exception as e:
|
||||
st.write("Unable to fetch data from Snowflake: ",e)
|
||||
# fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True)
|
||||
# fields = fields[~fields['SF_COL_NAME'].str.endswith('_PG', na=None)]
|
||||
|
||||
# change the code below if contract list is fetched from snowflake
|
||||
s3_client = boto3.client('s3',
|
||||
region_name="us-east-2"
|
||||
)
|
||||
|
||||
client_list, s3_paths = get_client_names()
|
||||
client_s3_paths = dict(zip(client_list, s3_paths))
|
||||
|
||||
client_row = st.columns([0.2, 0.7, 0.1])
|
||||
with client_row[0]:
|
||||
st.write("**Client Name**")
|
||||
with client_row[1]:
|
||||
client = st.selectbox('Client Name',(client_list), label_visibility = "collapsed", index= None)
|
||||
|
||||
# client_bucket = client_s3_paths.get(client)
|
||||
client_bucket = 'doczy-dev-infra-textract'
|
||||
|
||||
batch_objects = s3_client.list_objects_v2(Bucket=client_bucket
|
||||
, Prefix="contracts-landing-zone/", Delimiter='/')
|
||||
|
||||
batch_list = []
|
||||
for prefix in batch_objects['CommonPrefixes']:
|
||||
batch_list.append(prefix['Prefix'][:-1].split('/')[-1])
|
||||
|
||||
path_row = st.columns([0.1, 0.8])
|
||||
with path_row[0]:
|
||||
st.write("**Batch ID**")
|
||||
with path_row[1]:
|
||||
batch_id = st.selectbox('**Batch ID**', batch_list, label_visibility = "collapsed", index = None)
|
||||
|
||||
if batch_id:
|
||||
|
||||
objects = s3_client.list_objects_v2(Bucket=client_bucket, Prefix="contract-text-file/"+batch_id+"/")
|
||||
file_list = []
|
||||
for obj in objects['Contents']:
|
||||
if not obj['Key'].endswith('/'):
|
||||
file_list.append(obj['Key'])
|
||||
|
||||
contract_list = sorted(file_list)
|
||||
|
||||
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.selectbox('Select a file', ['All'] + contract_list, label_visibility = "collapsed", index= None)
|
||||
|
||||
field_row = st.columns([0.2, 0.7, 0.1])
|
||||
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", index = None)
|
||||
|
||||
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']
|
||||
|
||||
if st.button("Show Results"):
|
||||
query = 'select * from "DOCZY_PIPELINE_RAW_OUTPUT"'
|
||||
cur.execute(query)
|
||||
df2 = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
|
||||
|
||||
# get this dataframe from snowflake table
|
||||
# df2 = pd.DataFrame(columns=['Contract Name','Field Name', 'SF_DB_COL_NAME', 'Snippet','Page Number'
|
||||
# , 'Field Extracted Value', 'Actual Value','Imputed Value'])
|
||||
df2.to_csv('temp2.csv', index=False)
|
||||
|
||||
if st.button("Show PDF"):
|
||||
if file_name == None or file_name == "All":
|
||||
st.error("Choose one specific file.")
|
||||
else:
|
||||
with st.sidebar:
|
||||
st.markdown(
|
||||
"""
|
||||
<style>
|
||||
section[data-testid="stSidebar"] {
|
||||
width: 550px !important; # Set the width to your desired value
|
||||
}
|
||||
</style>
|
||||
""",
|
||||
unsafe_allow_html=True,
|
||||
)
|
||||
s3_obj = s3_client.get_object(Bucket = bucket, Key = 'textract-receiver-processed-pdfs/batch_2/2000-01-01 UCSD Medical Center PPA 14007097.PDF')
|
||||
data=s3_obj['Body'].read()
|
||||
pdf_viewer(data, width=1500)
|
||||
|
||||
# if st.button("Show PDF"):
|
||||
# if file_name == None or file_name == "All":
|
||||
# st.error("Choose one specific file.")
|
||||
# else:
|
||||
# with st.sidebar:
|
||||
# with open(file_name, "rb") as f:
|
||||
# base64_pdf = base64.b64encode(f.read()).decode('utf-8')
|
||||
# # Embedding PDF in HTML
|
||||
# pdf_display = F'<iframe src="data:application/pdf;base64,{base64_pdf}" width="500" height="1000" type="application/pdf"></iframe>'
|
||||
# # Displaying File
|
||||
# st.markdown(
|
||||
# """
|
||||
# <style>
|
||||
# section[data-testid="stSidebar"] {
|
||||
# width: 600px !important; # Set the width to your desired value
|
||||
# }
|
||||
# </style>
|
||||
# """,
|
||||
# unsafe_allow_html=True,
|
||||
# )
|
||||
# st.markdown(pdf_display, unsafe_allow_html=True)
|
||||
|
||||
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")
|
||||
st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
|
||||
with buttons[1]:
|
||||
# st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv')
|
||||
st.write("")
|
||||
with buttons[2]:
|
||||
if st.button("Kickoff Database Integration"):
|
||||
st.write("Stored in DB")
|
||||
|
||||
|
||||
|
||||
|
||||
@@ -0,0 +1,759 @@
|
||||
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
|
||||
import numpy as np
|
||||
from datetime import datetime
|
||||
import random
|
||||
import os
|
||||
import dateutil
|
||||
import util
|
||||
import anthropic
|
||||
import re
|
||||
import snowflake.connector
|
||||
from sf_conn import get_secret, save_to_sf
|
||||
from io import StringIO
|
||||
|
||||
|
||||
(REDIRECT_URI, create_batch_url, doczy_pipeline) = util.load_page_details(3)
|
||||
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")
|
||||
|
||||
_,c1= st.columns([5,1])
|
||||
# util.setup_page(REDIRECT_URI)
|
||||
try:
|
||||
util.setup_page(REDIRECT_URI)
|
||||
except:
|
||||
st.write("SSO Failed")
|
||||
st.session_state['user_info'] = {'mail': 'maamseek@aarete.com', 'displayName': 'Mayank Aamseek'}
|
||||
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()
|
||||
|
||||
try:
|
||||
sf_secrets = json.loads(get_secret())
|
||||
conn = snowflake.connector.connect(
|
||||
user=sf_secrets.get('user'),
|
||||
password=sf_secrets.get('password'),
|
||||
account="aarete-doczyai",
|
||||
role = "DEVADMIN",
|
||||
warehouse="DEV_XS",
|
||||
database="DOCZY_DEV",
|
||||
schema="STG"
|
||||
)
|
||||
cur = conn.cursor()
|
||||
query = 'select * from "TRAINING_DATA_RAW"'
|
||||
cur.execute(query)
|
||||
field_values = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
|
||||
# st.write(field_values)
|
||||
field_values['Document_Name'] = field_values['DOCUMENT_NAME']
|
||||
# field_values['Contract ID'] = field_values['CONTRACT_TITLE']
|
||||
# error('table values are incorrect')
|
||||
except:
|
||||
field_values = pd.read_csv('contract_field_values.csv', encoding='utf-8-sig', skipinitialspace=True)
|
||||
# field_values.rename(columns={'(internal) Document Name': 'Document_Name'}, inplace = True)
|
||||
## field_values.rename(columns={'(Internal) Carveout ID': 'Contract ID'}, inplace = True)
|
||||
field_values['Document_Name'] = field_values['DOCUMENT_NAME']
|
||||
field_values = field_values.loc[:, ~field_values.columns.str.contains('Unnamed:')]
|
||||
st.write("Local copy of TRAINING_DATA_RAW table loaded")
|
||||
|
||||
try:
|
||||
# error('table is not updated')
|
||||
query = 'select * from "BUSINESS_CONFIG"'
|
||||
cur.execute(query)
|
||||
fields = pd.DataFrame.from_records(iter(cur), columns=[x[0] for x in cur.description])
|
||||
fields.rename(columns={'FIELD_NAME': 'Field Name'}, inplace = True)
|
||||
fields.rename(columns={'QUESTION': 'Interrogation Question?'}, inplace = True)
|
||||
fields.rename(columns={'SF_COL_NAME': 'SF_DB_COL_NAME'}, inplace = True)
|
||||
fields = fields[~fields['SF_DB_COL_NAME'].str.endswith('_PG', na=None)]
|
||||
fields['Field Name'] = fields['SF_DB_COL_NAME']
|
||||
# fields.rename(columns={'FM_MODEL_ID': 'llm_selected'}, inplace = True)
|
||||
# error('table is empty')
|
||||
except:
|
||||
fields = pd.read_csv('contract_fields.csv', encoding='utf-8-sig', skipinitialspace=True)
|
||||
# fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name')
|
||||
# fields['Field Name'] = fields['SF_DB_COL_NAME']
|
||||
# fields = fields[~fields['Field Name'].isnull()]
|
||||
fields.rename(columns={'FIELD_NAME': 'Field Name'}, inplace = True)
|
||||
fields.rename(columns={'QUESTION': 'Interrogation Question?'}, inplace = True)
|
||||
fields.rename(columns={'SF_COL_NAME': 'SF_DB_COL_NAME'}, inplace = True)
|
||||
fields = fields[~fields['SF_DB_COL_NAME'].str.endswith('_PG', na=None)]
|
||||
fields['Field Name'] = fields['SF_DB_COL_NAME']
|
||||
st.write("Local copy of BUSINESS_CONFIG table loaded")
|
||||
|
||||
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)
|
||||
history = history[['Field Name', '# Contracts Tested', 'Username', 'Date/Time', 'Accuracy', 'Attempt #']]
|
||||
except:
|
||||
history = pd.read_csv('history.csv')
|
||||
st.write("Local copy of TRAINING_ATTEMPT_LOGS table loaded")
|
||||
|
||||
|
||||
|
||||
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 and Contract Related', 'Pricing Before Carveouts - All'
|
||||
, '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")
|
||||
|
||||
|
||||
# priorty column will be relaced by group_id in snowflake db
|
||||
if field_group == 'Unique and Contract Related':
|
||||
fields = fields[fields['PRIORITY'].isin(['A','C'])]
|
||||
elif field_group == 'Pricing Before Carveouts - All':
|
||||
fields = fields[fields['PRIORITY'] == 'B']
|
||||
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(' ')
|
||||
field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?']))
|
||||
|
||||
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 - Non Empty values', 'Multiple fields', 'One-to-many fields'), index=0, label_visibility = "collapsed")
|
||||
|
||||
field_row = st.columns([0.15, 0.45, 0.4])
|
||||
with field_row[0]:
|
||||
st.write("**Field Name**")
|
||||
with field_row[1]:
|
||||
if mode == 'Single field - Non Empty values':
|
||||
field = st.selectbox('Field Name',sorted(set(field_prompt_mapping.keys())), index=0, label_visibility = "collapsed")
|
||||
else:
|
||||
field = st.multiselect('Field Name',sorted(set(field_prompt_mapping.keys())), sorted(
|
||||
set(field_prompt_mapping.keys())), label_visibility = "collapsed")
|
||||
field_prompt_mapping = {key: field_prompt_mapping[key] for key in field}
|
||||
|
||||
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', '100', '200','Part-1', 'Part-2', 'All'), index=1, label_visibility = "collapsed")
|
||||
|
||||
s3_client = boto3.client('s3',
|
||||
region_name="us-east-2"
|
||||
)
|
||||
bucket = 'doczy-dev-infra-textract'
|
||||
objects = s3_client.list_objects_v2(Bucket=bucket, Prefix="training-data/")
|
||||
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))
|
||||
|
||||
if mode == 'Single field - 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()]
|
||||
|
||||
# df = pd.DataFrame({'col':contract_list})
|
||||
# st.write(df)
|
||||
# st.write(len(contract_list))
|
||||
# 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 = [str(contract)[:-4]+'.pdf' for contract in contract_list]
|
||||
contract_list = [contract for contract in contract_list if contract.rsplit('/',1)[1] in list(field_values['Document_Name'])]
|
||||
# st.write(len(contract_list))
|
||||
|
||||
if contract_count == 'All':
|
||||
contract_count = len(contract_list)
|
||||
if contract_count == 'Part-1':
|
||||
contract_list = np.array_split(contract_list, 2)[0]
|
||||
contract_count = len(contract_list)
|
||||
if contract_count == 'Part-2':
|
||||
contract_list = np.array_split(contract_list, 2)[1]
|
||||
contract_count = len(contract_list)
|
||||
|
||||
seed_row = st.columns([0.15, 0.45, 0.4])
|
||||
with seed_row[0]:
|
||||
if contract_count in ['10', '20', '30', '50', '100', '200']:
|
||||
st.write("**Seed Value**")
|
||||
elif contract_count == '1':
|
||||
st.write("**Contract Name**")
|
||||
with seed_row[1]:
|
||||
if contract_count in ['10', '20', '30', '50', '100', '200']:
|
||||
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)))
|
||||
contract_list = sorted(random.choices(contract_list, 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 3 - Haiku', 'Claude 3 - Sonnet', 'Claude Instant'
|
||||
, 'Llama 2 Chat 70B', 'Titan Text Express'), index=3, label_visibility = "collapsed")
|
||||
|
||||
st.write("**Prompt**")
|
||||
if mode != 'Single field - Non Empty values':
|
||||
sequence_input = json.dumps(field_prompt_mapping)
|
||||
else:
|
||||
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 = 100, label_visibility = "collapsed")
|
||||
|
||||
|
||||
# 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)
|
||||
if mode != 'One-to-many fields':
|
||||
field_values = field_values.drop_duplicates(subset='Contract Name', keep="first").sort_values('Contract Name')
|
||||
field_values = field_values[field_values['Contract Name'].isin([contract.rsplit('/',1)[1] for contract in contract_list])]
|
||||
|
||||
# Setup bedrock
|
||||
bedrock_runtime = boto3.client(
|
||||
service_name="bedrock-runtime",
|
||||
region_name="us-east-1",
|
||||
)
|
||||
|
||||
# question = prompt
|
||||
if mode != 'Single field - Non Empty values':
|
||||
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)
|
||||
question_with_schema = question
|
||||
attempt = 0
|
||||
|
||||
def run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, history):
|
||||
|
||||
# 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'])
|
||||
|
||||
field_list = []
|
||||
answer_list = []
|
||||
snippet_list = []
|
||||
page_no_list = []
|
||||
contract_list_f = []
|
||||
attempt = attempt + 1
|
||||
|
||||
for contract in contract_list:
|
||||
# with open(os.path.join(SOURCE_DIRECTORY, contract[:-4]+'.txt'), 'r') as infile:
|
||||
# context = infile.read()
|
||||
data = s3_client.get_object(Bucket=bucket, Key=str(contract)[:-4]+'.txt')
|
||||
contents = data['Body'].read()
|
||||
context = contents.decode("utf-8")
|
||||
# st.write(question_with_schema)
|
||||
|
||||
# Add "You must answer in correct JSON format."
|
||||
# Add Answer in JSON format: {{
|
||||
if llm_selected == "Titan Text Express":
|
||||
context = context[:16000]
|
||||
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.
|
||||
#
|
||||
{context}
|
||||
#
|
||||
|
||||
Question: {question}
|
||||
Answer: Answer in JSON format: {{"""
|
||||
parameters = {
|
||||
"maxTokenCount":2048,
|
||||
"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':
|
||||
context = context[:6000]
|
||||
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.
|
||||
##
|
||||
{context}
|
||||
##
|
||||
|
||||
Question: {question}
|
||||
Answer: Answer in JSON format: {{"""
|
||||
payload={
|
||||
"prompt":"[INST]"+ prompt_data +"[/INST]",
|
||||
"max_gen_len":2048,
|
||||
"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', 'Claude 3 - Haiku', 'Claude 3 - Sonnet']:
|
||||
if llm_selected == 'Claude Instant':
|
||||
context = context[:175000]
|
||||
|
||||
prompt_data = f"""
|
||||
|
||||
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.
|
||||
|
||||
{context}
|
||||
|
||||
Question: {question_with_schema}
|
||||
|
||||
Assistant: Answer in JSON format: {{"""
|
||||
|
||||
if llm_selected == "Claude 2":
|
||||
model_id = "anthropic.claude-v2:1"
|
||||
body = json.dumps(
|
||||
{"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT,
|
||||
"max_tokens_to_sample": 4096,
|
||||
"temperature":0.0,
|
||||
"top_p":1,
|
||||
"top_k":250,
|
||||
"stop_sequences":[anthropic.HUMAN_PROMPT]
|
||||
})
|
||||
elif llm_selected in ["Claude 3 - Haiku", "Claude 3 - Sonnet"]:
|
||||
if llm_selected == "Claude 3 - Haiku":
|
||||
model_id = 'anthropic.claude-3-haiku-20240307-v1:0'
|
||||
else:
|
||||
model_id = 'anthropic.claude-3-sonnet-20240229-v1:0'
|
||||
body = json.dumps({
|
||||
"anthropic_version": "bedrock-2023-05-31",
|
||||
"max_tokens": 4096,
|
||||
"messages": [
|
||||
{
|
||||
"role": "user",
|
||||
"content": [
|
||||
{
|
||||
"type": "text",
|
||||
"text":anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT
|
||||
}
|
||||
]
|
||||
}
|
||||
],
|
||||
"temperature": 0.0
|
||||
}
|
||||
)
|
||||
else:
|
||||
model_id = "anthropic.claude-instant-v1"
|
||||
body = json.dumps(
|
||||
{"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT,
|
||||
"max_tokens_to_sample": 2048,
|
||||
"temperature":0.0,
|
||||
"top_p":1,
|
||||
"top_k":250,
|
||||
"stop_sequences":[anthropic.HUMAN_PROMPT]
|
||||
})
|
||||
|
||||
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']
|
||||
elif llm_selected in ['Claude 3 - Haiku', 'Claude 3 - Sonnet']:
|
||||
response_text = response_body['content'][0]['text']
|
||||
|
||||
except:
|
||||
response_text = "failed"
|
||||
|
||||
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 + "}"
|
||||
|
||||
try:
|
||||
response_dict = json.loads(response_text)
|
||||
except:
|
||||
if mode != 'Single field - Non Empty values':
|
||||
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, " ")
|
||||
# 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]
|
||||
if mode != 'Single field - Non Empty values':
|
||||
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]
|
||||
else:
|
||||
field_l = [field]
|
||||
answer_l = [list(response_dict.values())[0]]
|
||||
try:
|
||||
page_no_l = [list(response_dict.values())[1]]
|
||||
except:
|
||||
page_no_l = ['']
|
||||
|
||||
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 ""
|
||||
|
||||
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]
|
||||
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 = [str(answer).strip("\n").strip().strip("[").strip("]").strip("{").strip("}").strip('"').rstrip('"').strip(
|
||||
' ') if answer is not None else None 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 3 - Haiku', 'Claude 3 - Sonnet', 'Claude Instant']:
|
||||
answer_list = [answer if "do not have" 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 "does not specify" 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 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 "Does not apply" not in str(answer) else " " for answer in answer_list]
|
||||
answer_list = [answer if "Nothing found" 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 Exception as e:
|
||||
st.write("post processing error")
|
||||
|
||||
# to be deleted later
|
||||
# contract_list_f = [contract.rsplit('/',1)[1].replace(' MU','').replace('_MU','').replace('.txt','') for contract in contract_list_f]
|
||||
contract_list_f = [contract.rsplit('/',1)[1] for contract in contract_list_f]
|
||||
|
||||
df['Contract Name'] = contract_list_f
|
||||
df['New Extracted value'] = answer_list
|
||||
df['Confidence Level'] = ' '
|
||||
df['Snippet'] = snippet_list
|
||||
df['New Page Number'] = page_no_list
|
||||
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 " ")
|
||||
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', 'LOB_Product_Network_Metal_Area_Program_Type_Speciality', 'SF_DB_COL_NAME'
|
||||
, 'Actual Value Stored', 'Original Page Number'])
|
||||
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]
|
||||
document_name = [x for x in list(field_values['Contract Name']) if not pd.isna(x) and file_name.rsplit('/',1)[1] in x][0]
|
||||
|
||||
temp_df = field_values[field_values['Contract Name'] == document_name].fillna('NA')
|
||||
unique_identifier = str(temp_df.at[temp_df.index[0],'CONTRACT_LOB']) + '__' + str(
|
||||
temp_df.at[temp_df.index[0],'CONTRACT_PRODUCT']) + '__' + str(
|
||||
temp_df.at[temp_df.index[0],'CONTRACT_NETWORK']) + '__' + str(
|
||||
temp_df.at[temp_df.index[0],'CONTRACT_MARKETPLACE_METAL_LEVEL']) + '__' + str(
|
||||
temp_df.at[temp_df.index[0],'CONTRACT_SERVICE_AREA']) + '__' + str(
|
||||
temp_df.at[temp_df.index[0],'CONTRACT_PROGRAM']) + '__' + str(
|
||||
temp_df.at[temp_df.index[0],'PROV_TYPE']) + '__' + str(temp_df.at[temp_df.index[0],'PROV_SPECIALTY'])
|
||||
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_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'])
|
||||
field_values_1['Contract Name'] = document_name
|
||||
field_values_1['LOB_Product_Network_Metal_Area_Program_Type_Speciality'] = unique_identifier
|
||||
field_values_2 = pd.concat([field_values_2, field_values_1], ignore_index = True)
|
||||
|
||||
if mode == 'One-to-many fields':
|
||||
for i in range(1, field_values[field_values['Contract Name'] == document_name].shape[0]):
|
||||
unique_identifier = str(temp_df.at[temp_df.index[i],'CONTRACT_LOB']) + '__' + str(
|
||||
temp_df.at[temp_df.index[i],'CONTRACT_PRODUCT']) + '__' + str(
|
||||
temp_df.at[temp_df.index[i],'CONTRACT_NETWORK']) + '__' + str(
|
||||
temp_df.at[temp_df.index[i],'CONTRACT_MARKETPLACE_METAL_LEVEL']) + '__' + str(
|
||||
temp_df.at[temp_df.index[i],'CONTRACT_SERVICE_AREA']) + '__' + str(
|
||||
temp_df.at[temp_df.index[i],'CONTRACT_PROGRAM']) + '__' + str(
|
||||
temp_df.at[temp_df.index[i],'PROV_TYPE']) + '__' + str(temp_df.at[temp_df.index[i],'PROV_SPECIALTY'])
|
||||
field_values_1 = field_values[field_values['Contract Name'] == document_name].iloc[[i]].transpose().reset_index()
|
||||
field_values_1.columns = ['SF_DB_COL_NAME', 'Actual Value Stored']
|
||||
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'])
|
||||
field_values_1['Contract Name'] = document_name
|
||||
field_values_1['LOB_Product_Network_Metal_Area_Program_Type_Speciality'] = unique_identifier
|
||||
field_values_2 = pd.concat([field_values_2, field_values_1], ignore_index = True)
|
||||
|
||||
|
||||
# st.write(df)
|
||||
# st.write(field_values_2)
|
||||
df = pd.merge(df, field_values_2, how ='right', on =['Contract Name', 'SF_DB_COL_NAME'])
|
||||
if mode != 'Single field - Non Empty values':
|
||||
df = df[df['SF_DB_COL_NAME'].isin(list(field_prompt_mapping.keys()))]
|
||||
else:
|
||||
df = df[df['SF_DB_COL_NAME'].isin([field])]
|
||||
if mode != 'One-to-many fields':
|
||||
df = df.drop_duplicates(subset=['SF_DB_COL_NAME', 'Contract Name'], keep="first")
|
||||
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 " ")
|
||||
df['Raw value 2'] = df['New Extracted value']
|
||||
df_date = df[df['SF_DB_COL_NAME'].str.contains('_DT', na=False)]
|
||||
df_others = df[~df['SF_DB_COL_NAME'].str.contains('_DT', na=False)]
|
||||
|
||||
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)
|
||||
|
||||
df.fillna(" ", inplace=True)
|
||||
df['Raw value 3'] = df['New Extracted value']
|
||||
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 '')
|
||||
actual_value_list = list(df['Actual Value Stored'])
|
||||
actual_value_list = [answer if str(answer) != "12 months" else "1 year" for answer in actual_value_list]
|
||||
actual_value_list = [answer if str(answer) != "Fifth" else "5" for answer in actual_value_list]
|
||||
actual_value_list = [answer if str(answer) != "Seventh" else "7" for answer in actual_value_list]
|
||||
|
||||
answer_list = list(df['New Extracted value'])
|
||||
answer_list = [answer if str(answer) != "one-year" else "1 year" for answer in answer_list]
|
||||
answer_list = [answer if str(answer) != "one year" else "1 year" for answer in answer_list]
|
||||
answer_list = [answer if str(answer) != "one" else "1 year" for answer in answer_list]
|
||||
answer_list = [answer if str(answer) != "one (1) year" else "1 year" for answer in answer_list]
|
||||
answer_list = [answer if str(answer) != "twelve" else "1 year" for answer in answer_list]
|
||||
answer_list = [answer if str(answer) != "XI" else "11" for answer in answer_list]
|
||||
answer_list = [answer if str(answer) != "Third" else "3" for answer in answer_list]
|
||||
answer_list = [answer if str(answer) != "Six" else "6" for answer in answer_list]
|
||||
|
||||
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)]
|
||||
df['Result'] = [str(x) for x in result_list]
|
||||
|
||||
# df = df[~df['Contract ID'].isnull()]
|
||||
# df['Contract ID'] = contract_list_f
|
||||
df['Contract ID'] = range(len(actual_value_list))
|
||||
|
||||
df = df[['Contract Name','Contract ID', 'LOB_Product_Network_Metal_Area_Program_Type_Speciality', 'SF_DB_COL_NAME', 'Actual Value Stored', 'Raw value', 'Raw value 2', 'Raw value 3'
|
||||
, 'New Extracted value','Confidence Level','Snippet','Original Page Number', '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'
|
||||
|
||||
if mode != 'Single field - Non Empty values':
|
||||
field = field_group
|
||||
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]
|
||||
# df.to_csv("RESULTS\\"+field.replace("?","").replace("/","_")+'-'+llm_selected+'.csv', index=False)
|
||||
|
||||
return df, history, attempt, raw_response_text
|
||||
|
||||
raw_response_text = ''
|
||||
df = pd.DataFrame(columns=['Contract Name','Contract ID', 'Actual Value Stored','New Extracted value','Confidence Level'
|
||||
,'Snippet','New Page Number', 'Revised Prompt', 'Result'])
|
||||
|
||||
if st.button("Test Configuration"):
|
||||
df, history, attempt, raw_response_text = run_llm(attempt, bucket, contract_list, llm_selected, field, field_values, history)
|
||||
df_1 = df[['Contract Name','Contract ID', 'Actual Value Stored','New Extracted value','Confidence Level'
|
||||
,'Snippet','Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
|
||||
df.to_csv('results.csv', index=False)
|
||||
history.to_csv('history.csv', index=False)
|
||||
# s3_client.upload_file('results.csv', bucket, key)
|
||||
csv_buf = StringIO()
|
||||
df_1.to_csv(csv_buf, header=True, index=False)
|
||||
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()
|
||||
history.tail(1).to_csv(csv_buf, header=True, index=False)
|
||||
csv_buf.seek(0)
|
||||
s3_client.put_object(Bucket='doczy-dev-infra-raw-data-ingestion', Body=csv_buf.getvalue(), Key='training_interface/history.csv')
|
||||
df = df[['Contract Name','Contract ID', 'LOB_Product_Network_Metal_Area_Program_Type_Speciality','SF_DB_COL_NAME', 'Actual Value Stored','New Extracted value','Confidence Level'
|
||||
,'Snippet','Original Page Number', 'New Page Number', 'Revised Prompt', 'Result']]
|
||||
|
||||
try:
|
||||
df = pd.read_csv('results.csv')
|
||||
df['Result'] = df['Result'].astype('str')
|
||||
except:
|
||||
df = pd.DataFrame(columns=['Contract Name' ,'New Extracted value','Confidence Level','Snippet','New Page Number'
|
||||
, 'Revised Prompt', 'Result'])
|
||||
history = pd.read_csv('history.csv')
|
||||
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")
|
||||
|
||||
add_vertical_space(20)
|
||||
st.write(field)
|
||||
# if mode != 'Multiple fields':
|
||||
# st.write(fields.loc[fields['SF_DB_COL_NAME'] == field, 'Field Name'].iloc[0])
|
||||
|
||||
st.write(len(contract_list))
|
||||
st.write(raw_response_text)
|
||||
|
||||
try:
|
||||
save_to_sf('load_training_results', training_results_file_name = "results.csv", attempt_logs_file_name = "history.csv")
|
||||
except:
|
||||
st.write("running locally")
|
||||
|
||||
|
||||
@@ -0,0 +1,77 @@
|
||||
import streamlit as st
|
||||
import msal
|
||||
import requests
|
||||
import boto3
|
||||
from botocore.exceptions import ClientError
|
||||
import json
|
||||
|
||||
# Replace with your own values
|
||||
CLIENT_ID = 'effafe90-7ed7-43a3-ab03-19a0be2f1758'
|
||||
# CLIENT_SECRET = 'bjQ8Q~lpR2uBcGI34VDu16t73doz8Crj0YY_~dgD'
|
||||
# TENANT_ID = ''
|
||||
|
||||
AUTHORITY = 'https://login.microsoftonline.com/organizations/'
|
||||
SCOPE = ['User.Read']
|
||||
REDIRECT_URI = 'https://172.29.20.102:8501'
|
||||
|
||||
# Initialize boto3 client to interact with AWS Secrets Manager
|
||||
|
||||
|
||||
|
||||
|
||||
def get_secret():
|
||||
|
||||
secret_name = "doczy-sso-azure-app-key"
|
||||
region_name = "us-east-2"
|
||||
|
||||
# Create a Secrets Manager client
|
||||
# session = boto3.session.Session()
|
||||
# client = session.client(
|
||||
# service_name='secretsmanager',
|
||||
# region_name=region_name
|
||||
# )
|
||||
client = boto3.client('secretsmanager', region_name=region_name)
|
||||
try:
|
||||
get_secret_value_response = client.get_secret_value(
|
||||
SecretId=secret_name
|
||||
)
|
||||
except ClientError as e:
|
||||
raise e
|
||||
|
||||
secret = get_secret_value_response['SecretString']
|
||||
secret = json.loads(secret)['CLIENT_SECRET']
|
||||
|
||||
return secret
|
||||
|
||||
|
||||
CLIENT_SECRET = get_secret()
|
||||
|
||||
app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET)
|
||||
|
||||
|
||||
def get_auth_url(REDIRECT_URI):
|
||||
auth_url = app.get_authorization_request_url(SCOPE, redirect_uri=REDIRECT_URI)
|
||||
return auth_url
|
||||
|
||||
@st.cache_data
|
||||
def get_token_from_code(auth_code, REDIRECT_URI):
|
||||
app = msal.ConfidentialClientApplication(CLIENT_ID, authority=AUTHORITY, client_credential=CLIENT_SECRET)
|
||||
result = app.acquire_token_by_authorization_code(auth_code, scopes=SCOPE, redirect_uri=REDIRECT_URI)
|
||||
return result['access_token']
|
||||
|
||||
|
||||
def get_user_info(access_token):
|
||||
headers = {'Authorization': f'Bearer {access_token}'}
|
||||
response = requests.get('https://graph.microsoft.com/v1.0/me', headers=headers)
|
||||
return response.json()
|
||||
|
||||
|
||||
def handle_redirect(REDIRECT_URI):
|
||||
if not st.session_state.get('access_token'):
|
||||
code = st.query_params.get('code')
|
||||
if code:
|
||||
access_token = get_token_from_code(code, REDIRECT_URI)
|
||||
st.session_state['access_token'] = access_token
|
||||
st.session_state
|
||||
|
||||
|
||||
@@ -0,0 +1,140 @@
|
||||
# Use this code snippet in your app.
|
||||
# If you need more information about configurations
|
||||
# or implementing the sample code, visit the AWS docs:
|
||||
# https://aws.amazon.com/developer/language/python/
|
||||
|
||||
import boto3
|
||||
from botocore.exceptions import ClientError
|
||||
import http.client
|
||||
import base64
|
||||
import ast
|
||||
import snowflake.connector
|
||||
import json
|
||||
|
||||
|
||||
def get_secret():
|
||||
|
||||
secret_name = "doczy_dev_db_creds"
|
||||
region_name = "us-east-2"
|
||||
|
||||
# Create a Secrets Manager client
|
||||
session = boto3.session.Session()
|
||||
client = session.client(
|
||||
service_name='secretsmanager',
|
||||
region_name=region_name
|
||||
)
|
||||
|
||||
try:
|
||||
get_secret_value_response = client.get_secret_value(
|
||||
SecretId=secret_name
|
||||
)
|
||||
except ClientError as e:
|
||||
# For a list of exceptions thrown, see
|
||||
# https://docs.aws.amazon.com/secretsmanager/latest/apireference/API_GetSecretValue.html
|
||||
raise e
|
||||
|
||||
secret = get_secret_value_response['SecretString']
|
||||
return secret
|
||||
|
||||
# get_secret()
|
||||
|
||||
# TODO: This function needs to be changed to accept Kwargs
|
||||
# The function name should be more generic and cofnigurable
|
||||
def save_to_sf(dag_name, **kwargs):
|
||||
|
||||
mwaa_env_name = 'doczy-dev-infra-mwaa'
|
||||
dag_name = dag_name
|
||||
mwaa_cli_command = 'dags trigger'
|
||||
|
||||
# Create the client with the specified profile
|
||||
session = boto3.Session()
|
||||
client = session.client('mwaa', region_name='us-east-2')
|
||||
|
||||
# get web token
|
||||
mwaa_cli_token = client.create_cli_token(
|
||||
Name=mwaa_env_name
|
||||
)
|
||||
|
||||
conn = http.client.HTTPSConnection(mwaa_cli_token['WebServerHostname'])
|
||||
|
||||
# This section passes the payload to the MWAA CLI
|
||||
# The file parameters should be added dynamically in streamlit, once the file names are passed while triggering the dag, the data will be ingested
|
||||
# training_results_file = "training_results_sample.csv"
|
||||
# attempt_logs_file = "attempt_logs_sample.csv"
|
||||
# conf = "{\"" + "training_results_file_name" + "\":\"" + {training_results_file} + "\", \"" + "attempt_logs_file_name" + "\":\"" + {attempt_logs_file} + "\"}".format(training_results_file=training_results_file, attempt_logs_file=attempt_logs_file)
|
||||
|
||||
conf = json.dumps(kwargs)
|
||||
|
||||
payload = mwaa_cli_command + " " + dag_name + " --conf '{}'".format(conf)
|
||||
headers = {
|
||||
'Authorization': 'Bearer ' + mwaa_cli_token['CliToken'],
|
||||
'Content-Type': 'text/plain'
|
||||
}
|
||||
conn.request("POST", "/aws_mwaa/cli", payload, headers)
|
||||
res = conn.getresponse()
|
||||
data = res.read()
|
||||
dict_str = data.decode("UTF-8")
|
||||
mydata = ast.literal_eval(dict_str)
|
||||
return payload
|
||||
|
||||
|
||||
# save_to_sf("2024-03-13T17-36_prompt_results.csv", "2024-03-14T23-19_history.csv")
|
||||
|
||||
|
||||
# Create snowflake connection with the secrets fetched from secrets manager for the schema passed as an argument
|
||||
def get_snowflake_conn(schema):
|
||||
secret = get_secret()
|
||||
secret_dict = eval(secret)
|
||||
conn = snowflake.connector.connect(
|
||||
user=secret_dict['user'],
|
||||
password=secret_dict['password'],
|
||||
account=secret_dict['account_alias'],
|
||||
warehouse=secret_dict['warehouse'],
|
||||
database=secret_dict['database'],
|
||||
schema=schema
|
||||
)
|
||||
return conn
|
||||
|
||||
def get_client_names():
|
||||
"""
|
||||
Input: None
|
||||
Output: client_names, s3_paths lists
|
||||
"""
|
||||
try:
|
||||
# Get conn from snowflake_conn function for STG schema
|
||||
conn = get_snowflake_conn('STG')
|
||||
cursor = conn.cursor()
|
||||
|
||||
# Query to get client names and their s3_paths
|
||||
query = "SELECT DISTINCT client_name, bucket_name FROM STG.CLIENT_LOGS"
|
||||
cursor.execute(query)
|
||||
# Create 2 lists from the query results
|
||||
client_names = []
|
||||
s3_paths = []
|
||||
|
||||
for row in cursor:
|
||||
client_names.append(row[0])
|
||||
s3_paths.append(row[1]) # Extracting the bucket name from the s3 path
|
||||
cursor.close()
|
||||
conn.close()
|
||||
return client_names, s3_paths
|
||||
except Exception as e:
|
||||
print(f"Error while fetching client names from snowflake: {e}")
|
||||
return None, None
|
||||
|
||||
|
||||
def insert_upload_logs(batch_id, client_name, file_name, upload_datetime, upload_user):
|
||||
"""
|
||||
Input: batch_id, client_name, file_name, upload_datetime, upload_user
|
||||
Output: status of the insert query
|
||||
"""
|
||||
try:
|
||||
conn = get_snowflake_conn('STG')
|
||||
cursor = conn.cursor()
|
||||
query = f"INSERT INTO STG.CONTRACT_UPLOAD_LOGS (BATCH_ID, CLIENT_NAME, FILE_NAME, UPLOAD_DATETIME, UPLOAD_USER) VALUES ('{batch_id}', '{client_name}', '{file_name}', '{upload_datetime}', '{upload_user}')"
|
||||
cursor.execute(query)
|
||||
cursor.close()
|
||||
conn.close()
|
||||
return 'Log inserted successfully'
|
||||
except Exception as e:
|
||||
return e
|
||||
@@ -0,0 +1,35 @@
|
||||
|
||||
import streamlit as st
|
||||
import security
|
||||
import os
|
||||
import constants
|
||||
|
||||
def setup_page(REDIRECT_URI):
|
||||
# st.set_page_config(
|
||||
# page_title=page_title,
|
||||
# page_icon="👋",
|
||||
# )
|
||||
|
||||
if st.query_params.get('code'):
|
||||
security.handle_redirect(REDIRECT_URI)
|
||||
|
||||
access_token = st.session_state.get('access_token')
|
||||
|
||||
if access_token:
|
||||
user_info = security.get_user_info(access_token)
|
||||
st.session_state['user_info'] = user_info
|
||||
return True
|
||||
else:
|
||||
st.write("Please sign-in to use this app.")
|
||||
auth_url = security.get_auth_url(REDIRECT_URI)
|
||||
st.markdown(f"<a href='{auth_url}' target='_self'>Sign In</a>", unsafe_allow_html=True)
|
||||
st.stop()
|
||||
|
||||
def load_page_details(interface):
|
||||
env_var = os.environ.get('ENVIRONMENT', 'DEV')
|
||||
if env_var == 'UAT':
|
||||
return (constants.DOCZY_REDIRECT_URL_UAT + str(interface), constants.DOCZY_CREATE_BATCH_URL_UAT, constants.DOCZY_PIPELINE_URL_UAT)
|
||||
elif env_var == 'DEV':
|
||||
return (constants.DOCZY_REDIRECT_URL_DEV + str(interface), constants.DOCZY_CREATE_BATCH_URL_DEV, constants.DOCZY_PIPELINE_URL_DEV)
|
||||
else:
|
||||
return (constants.DOCZY_REDIRECT_URL_DEV + str(interface), constants.DOCZY_CREATE_BATCH_URL_DEV, constants.DOCZY_PIPELINE_URL_DEV)
|
||||
Reference in New Issue
Block a user