diff --git a/streamlit/interface_1.py b/streamlit/interface_1.py index d4398de..aaf3542 100644 --- a/streamlit/interface_1.py +++ b/streamlit/interface_1.py @@ -8,7 +8,7 @@ from datetime import datetime import boto3 import util -REDIRECT_URI = 'http://172.29.20.126:8501' +REDIRECT_URI = 'https://doczy.aarete.com:8501' user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' , 'piragavarapu@aarete.com', 'umistry@aarete.com'] st.set_page_config(layout = "wide") @@ -25,9 +25,15 @@ st.set_page_config(layout = "wide") # 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: +try: + util.setup_page(REDIRECT_URI) + _,c1= st.columns([5,1]) + c1.write(f"User: **{st.session_state.user_info['displayName']}**") + user_mail = st.session_state.user_info['mail'] +except: + user_mail = 'maamseek@aarete.com' + +if user_mail in user_list: s3_client = boto3.client('s3', region_name="us-east-1", @@ -39,11 +45,19 @@ if 'maamseek@aarete.com' in user_list: for prefix in objects['CommonPrefixes']: folder_list.append(prefix['Prefix'][:-1].split('/')[-1]) + # to be deleted later + folder_list = ['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_row = st.columns([0.1, 0.8]) with client_row[0]: st.write("**Client Name**") with client_row[1]: - client = st.selectbox('Client Name',(folder_list), index=7, label_visibility = "collapsed") + client = st.selectbox('Client Name',(folder_list), label_visibility = "collapsed") + + # to be deleted later + client = 'textract-receiver-processed-file' folder_objects = s3_client.list_objects_v2(Bucket='doczy-dev-infra-textract' , Prefix="batches/batch_1/"+client+"/", Delimiter='/') @@ -78,7 +92,7 @@ if 'maamseek@aarete.com' in user_list: add_vertical_space(1) - df = pd.DataFrame(columns=['Request ID','Contract ID','Contract Name','Unique Key','Pricing Before Carveouts' + df = pd.DataFrame(columns=['Request ID','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='doczy-dev-infra-textract' @@ -91,7 +105,7 @@ if 'maamseek@aarete.com' in user_list: df['Contract Name'] = file_list df['Request ID'] = range(len(file_list)) - df['Contract ID'] = file_list + # df['Contract ID'] = file_list df['Unique Key'] = a df['Pricing Before Carveouts'] = b df['Contract Related'] = c diff --git a/streamlit/interface_2.py b/streamlit/interface_2.py index f440ed8..00ab017 100644 --- a/streamlit/interface_2.py +++ b/streamlit/interface_2.py @@ -12,9 +12,14 @@ import streamlit as st from streamlit_extras.add_vertical_space import add_vertical_space import os import pandas as pd +import numpy as np import util +import anthropic +from pydantic import BaseModel +from typing import List +import re -REDIRECT_URI = 'http://localhost:8502' +REDIRECT_URI = 'https://doczy.aarete.com:8502' user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' , 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com' , 'vnair@aarete.com'] @@ -33,20 +38,22 @@ with st.sidebar: add_vertical_space(15) # st.write("Doczy") -util.setup_page(REDIRECT_URI) -if st.session_state.user_info['mail'] in user_list: - - fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) - fields = fields[fields['PRIORITY'] == 'A'] - fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name') - fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ') - field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?'])) +try: + util.setup_page(REDIRECT_URI) + _,c1= st.columns([4,1]) + c1.write(f"User: **{st.session_state.user_info['displayName']}**") + user_mail = st.session_state.user_info['mail'] +except: + user_mail = 'maamseek@aarete.com' - def file_selector(folder_path=SOURCE_DIRECTORY): - filenames = os.listdir(folder_path) - selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed") - # return os.path.join(folder_path, selected_filename) - return selected_filename +def file_selector(folder_path=SOURCE_DIRECTORY): + filenames = os.listdir(folder_path) + selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed") + # return os.path.join(folder_path, selected_filename) + return selected_filename + + +if user_mail in user_list: file_row = st.columns([0.2, 0.7, 0.1]) with file_row[0]: @@ -61,25 +68,69 @@ if st.session_state.user_info['mail'] in user_list: # with lob_row[1]: # lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed") + 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") + llm_row = st.columns([0.2, 0.7, 0.1]) with llm_row[0]: st.write("**Langauge Model**") with llm_row[1]: - llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude Instant', 'Llama 2 Chat 13B', 'Llama 2 Chat 70B' - , 'Titan Text Express'), label_visibility = "collapsed") + llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude Instant', 'Llama 2 Chat 70B' + , 'Titan Text Express'), index=1, label_visibility = "collapsed") + + fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) + fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name') + fields = fields[~fields['Field Name'].isnull()] + + 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'] + + fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ') + field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?'])) - page_list = [] with open(os.path.join(SOURCE_DIRECTORY, file_name), 'r') as infile: - text = infile.read() - page_count = text.count('Start of Page No. = ') - for page in range(page_count+1): - file_path = "SOURCE_DOCUMENTS\\" + f'{file_name[:-4]}_page{page}.txt' - dict_with_pages = { 'source': { '$eq': file_path }} - page_list.append(dict_with_pages) - - # AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID') - # AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY') - # AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN') + context = infile.read() + # page_count = context.count('Start of Page No. = ') # Setup bedrock bedrock_runtime = boto3.client( @@ -87,120 +138,203 @@ if st.session_state.user_info['mail'] in user_list: region_name="us-east-1" ) - embeddings = BedrockEmbeddings( - client=bedrock_runtime, - model_id="amazon.titan-embed-text-v1", - ) - DB = Chroma( - persist_directory=PERSIST_DIRECTORY, - embedding_function=embeddings, - client_settings=CHROMA_SETTINGS, - ) - RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4}) + # question_list = list(field_prompt_mapping.items()) + # class OutputSchema(BaseModel): + # question_list[0]: str + # question_list[1]: str + # question_list[2]: str + # question_list[3]: str + # question_list[4]: str + # question_list[5]: str + # question_list[6]: str + # question_list[7]: str + # question_list[8]: str + # question_list[9]: str + # question_list[10]: str + # question_list[11]: str + # question_list[12]: str + # question_list[13]: str + # question_list[14]: str + # question_list[15]: str + # question_list[16]: str + # question_list[17]: str + # question_list[18]: str + # question_list[19]: str - # 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, + # class OutputList(BaseModel): + # answer: List[OutputSchema] + + # question = '\\n'.join(question_list) + question = json.dumps(field_prompt_mapping) + # question_with_schema = f'{question}{OutputList.schema_json()}' + question_with_schema = question + + 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":512, + "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': - 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 + context = context[:7000] + 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} + ## - template = """ + Question: {question} + Answer: Answer in JSON format: {{ + """ + payload={ + "prompt":"[INST]"+ prompt_data +"[/INST]", + "max_gen_len":512, + "temperature":0.0, + "top_p":0.9 + } + body=json.dumps(payload) + model_id="meta.llama2-70b-chat-v1" - Use the following pieces of context to answer the question at the end. If you don't know the answer,\ - just say that you don't know, don't try to make up an answer. + elif llm_selected in ['Claude Instant', 'Claude 2']: - {context} + prompt_data = f""" - Question: {question} - Answer:""" - prompt = PromptTemplate(input_variables=["context", "question"], template=template) + 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 JSON format. - QA = RetrievalQA.from_chain_type( - llm=LLM, - chain_type="stuff", - retriever=RETRIEVER, - return_source_documents=True, - chain_type_kwargs={"prompt": prompt}, - ) + {context} - # query = "In which state or states is the Contract applicable? Answer in one or two words. State name: " - # response = QA({"query":query}) - # st.write(query) - # st.write(response['result']) - # st.write("-----------") - # st.write(response) + Question: {question_with_schema} - # clicked = st.button("Show Results") - df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level', + Assistant: Answer in JSON format: {{ + """ + body = json.dumps( + {"prompt": anthropic.HUMAN_PROMPT + prompt_data + anthropic.AI_PROMPT, + "max_tokens_to_sample": 1024, + "temperature":0.0, + "top_p":1, + "top_k":250, + "stop_sequences":[anthropic.HUMAN_PROMPT] + }) + if llm_selected == "Claude 2": + model_id = "anthropic.claude-v2:1" + else: + model_id = "anthropic.claude-instant-v1" + + + # def claude_prompt_format(prompt: str) -> str: + # # Add headers to start and end of prompt + # return "\n\nHuman: " + prompt + "\n\nAssistant:" + + # # Call Claude model + # def call_claude(prompt): + # prompt_config = { + # "prompt": claude_prompt_format(prompt), + # "max_tokens_to_sample": 4096, + # "temperature": 0.5, + # "top_k": 250, + # "top_p": 0.5, + # "stop_sequences": [], + # } + + # body = json.dumps(prompt_config) + + # modelId = "anthropic.claude-instant-v1" + # accept = "application/json" + # contentType = "application/json" + + # response = bedrock_runtime.invoke_model( + # body=body, modelId=modelId, accept=accept, contentType=contentType + # ) + # response_body = json.loads(response.get("body").read()) + + # results = response_body.get("completion") + # return results + + # prompt = SECOND_PROMPT + # result = call_claude(prompt) + # st.write(result) + + df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number', 'Field Extracted Value','Imputed Value']) - field_list = list(field_prompt_mapping.keys()) - query_list = [field_prompt_mapping[x] for x in field_list] - score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list] - confidence_list = [] - for score in score_list: - confidence_list.append(max(d[1] for d in score)) - # st.write(confidence_list) + # field_list = list(field_prompt_mapping.keys()) + # query_list = [field_prompt_mapping[x] for x in field_list] + # st.write(field_prompt_mapping) + # st.write(prompt_data) if st.button("Show Results"): - response_list = [QA({"query":query}) for query in query_list] - answer_list = [response['result'] for response in response_list] - doc_list = [response['source_documents'] for response in response_list] - snippet_list = [str(doc[0].page_content) for doc in doc_list] - page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list] + + 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'] + + # st.write(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: + response_dict = {"Test value": "Failed to extract"} + + # st.write(response_dict) + field_list = list(response_dict.keys()) + answer_list = list(response_dict.values()) + # st.write(answer_list) + location_list = [context.find(answer) if isinstance(answer, str) and answer != "" else -1 for answer in answer_list] + snippet_list = [' '.join(context[:location].split('.')[-4:]) + ' ' + ' '.join(context[location:].split('. ')[:5] + ) if location != -1 else ' ' for location in location_list] + page_no_list = [" " 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_list] + # st.write(location_list) + page_no_list = [re.search(r'\d+', page).group() if page != " " and re.search(r'\d+', page) is not None else "" for page in page_no_list] + + + # response_list = [QA({"query":query}) for query in query_list] + # answer_list = [response['result'] for response in response_list] + # doc_list = [response['source_documents'] for response in response_list] + # snippet_list = [str(doc[0].page_content) for doc in doc_list] + # page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list] df['Field Name'] = field_list df['Contract Name'] = file_name df['Snippet'] = snippet_list df['Page Number'] = page_no_list - df['Confidence Level'] = confidence_list + # df['Confidence Level'] = ' ' df['Field Extracted Value'] = answer_list df.to_csv('temp2.csv', index=False) diff --git a/streamlit/interface_2_rag.py b/streamlit/interface_2_rag.py new file mode 100644 index 0000000..3b4d8a6 --- /dev/null +++ b/streamlit/interface_2_rag.py @@ -0,0 +1,232 @@ +import json + +import boto3 +from langchain.prompts import PromptTemplate +from langchain.embeddings.bedrock import BedrockEmbeddings +from langchain.llms.bedrock import Bedrock +from langchain_community.vectorstores import Chroma +from constants import CHROMA_SETTINGS, EMBEDDING_MODEL_NAME, PERSIST_DIRECTORY, MODEL_ID, MODEL_BASENAME, SOURCE_DIRECTORY +from langchain.chains import RetrievalQA + +import streamlit as st +from streamlit_extras.add_vertical_space import add_vertical_space +import os +import pandas as pd +import util + +REDIRECT_URI = 'https://doczy.aarete.com:8502' +user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' +, 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com' +, 'vnair@aarete.com'] + +st.set_page_config(layout = "wide") +# Sidebar contents +with st.sidebar: + st.title("Doczy.AI ™") + st.markdown( + """ + ## About + This app extracts data from contracts + + """ + ) + add_vertical_space(15) + # st.write("Doczy") + +util.setup_page(REDIRECT_URI) +_,c1= st.columns([5,1]) +c1.write(f"User: **{st.session_state.user_info['displayName']}**") + +if st.session_state.user_info['mail'] in user_list: + + fields = pd.read_csv('contract_fields.csv', encoding='unicode_escape', skipinitialspace=True) + fields = fields[fields['PRIORITY'] == 'A'] + fields = fields.drop_duplicates(subset='Field Name', keep="first").sort_values('Field Name') + fields['Interrogation Question?'] = fields['Interrogation Question?'].fillna(' ') + field_prompt_mapping = dict(zip(fields['Field Name'], fields['Interrogation Question?'])) + + def file_selector(folder_path=SOURCE_DIRECTORY): + filenames = os.listdir(folder_path) + selected_filename = st.selectbox('Select a file', filenames, label_visibility = "collapsed") + # return os.path.join(folder_path, selected_filename) + return selected_filename + + file_row = st.columns([0.2, 0.7, 0.1]) + with file_row[0]: + st.write("**Contract Name**") + with file_row[1]: + # file_name = st.text_input("**Contract Name**", label_visibility = "collapsed") + file_name = file_selector() + + # lob_row = st.columns([0.2, 0.7, 0.1]) + # with lob_row[0]: + # st.write("**LOB**") + # with lob_row[1]: + # lob = st.selectbox('LOB',('Medicare', 'Medicaid'), label_visibility = "collapsed") + + llm_row = st.columns([0.2, 0.7, 0.1]) + with llm_row[0]: + st.write("**Langauge Model**") + with llm_row[1]: + llm_selected = st.selectbox('Langauge Model',('Claude 2', 'Claude Instant', 'Llama 2 Chat 13B', 'Llama 2 Chat 70B' + , 'Titan Text Express'), label_visibility = "collapsed") + + page_list = [] + with open(os.path.join(SOURCE_DIRECTORY, file_name), 'r') as infile: + text = infile.read() + page_count = text.count('Start of Page No. = ') + for page in range(page_count+1): + file_path = "SOURCE_DOCUMENTS\\" + f'{file_name[:-4]}_page{page}.txt' + dict_with_pages = { 'source': { '$eq': file_path }} + page_list.append(dict_with_pages) + + # AWS_ACCESS_KEY_ID = os.getenv('AWS_ACCESS_KEY_ID') + # AWS_SECRET_ACCESS_KEY = os.getenv('AWS_SECRET_ACCESS_KEY') + # AWS_SESSION_TOKEN=os.getenv('AWS_SESSION_TOKEN') + + # Setup bedrock + bedrock_runtime = boto3.client( + service_name="bedrock-runtime", + region_name="us-east-1" + ) + + embeddings = BedrockEmbeddings( + client=bedrock_runtime, + model_id="amazon.titan-embed-text-v1", + ) + DB = Chroma( + persist_directory=PERSIST_DIRECTORY, + embedding_function=embeddings, + client_settings=CHROMA_SETTINGS, + ) + RETRIEVER = DB.as_retriever(search_kwargs={"filter": {"$or": page_list}, "k": 4}) + + # if "LLM" 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 + + template = """ + + Use the following pieces of context to answer the question at the end. If you don't know the answer,\ + just say that you don't know, don't try to make up an answer. + + {context} + + Question: {question} + Answer:""" + prompt = PromptTemplate(input_variables=["context", "question"], template=template) + + QA = RetrievalQA.from_chain_type( + llm=LLM, + chain_type="stuff", + retriever=RETRIEVER, + return_source_documents=True, + chain_type_kwargs={"prompt": prompt}, + ) + + # query = "In which state or states is the Contract applicable? Answer in one or two words. State name: " + # response = QA({"query":query}) + # st.write(query) + # st.write(response['result']) + # st.write("-----------") + # st.write(response) + + # clicked = st.button("Show Results") + df = pd.DataFrame(columns=['Contract Name','Field Name','Snippet','Page Number','Confidence Level', + 'Field Extracted Value','Imputed Value']) + field_list = list(field_prompt_mapping.keys()) + query_list = [field_prompt_mapping[x] for x in field_list] + score_list = [DB.similarity_search_with_relevance_scores(query, k=4, filter={"$or": page_list}) for query in query_list] + confidence_list = [] + for score in score_list: + confidence_list.append(max(d[1] for d in score)) + # st.write(confidence_list) + + if st.button("Show Results"): + response_list = [QA({"query":query}) for query in query_list] + answer_list = [response['result'] for response in response_list] + doc_list = [response['source_documents'] for response in response_list] + snippet_list = [str(doc[0].page_content) for doc in doc_list] + page_no_list = [int(str(doc[0].metadata["source"]).rsplit('_page')[1].replace('.txt',''))+1 for doc in doc_list] + + df['Field Name'] = field_list + df['Contract Name'] = file_name + df['Snippet'] = snippet_list + df['Page Number'] = page_no_list + df['Confidence Level'] = confidence_list + df['Field Extracted Value'] = answer_list + df.to_csv('temp2.csv', index=False) + + df2 = pd.read_csv('temp2.csv') + df2['Imputed Value'] = '' + edited_df = st.data_editor(df2) + + @st.cache_data + def convert_df(df): + return df.to_csv(index=False).encode('utf-8') + + csv = convert_df(edited_df) + + buttons = st.columns(3) + with buttons[0]: + st.button("Save All Imputations") + with buttons[1]: + st.download_button("Download Table", csv, "file.csv", "text/csv", key='download-csv') + with buttons[2]: + st.button("Kickoff Database Integration") + +else: + st.write("Access Denied") + + + diff --git a/streamlit/interface_3.py b/streamlit/interface_3.py index c9ecfae..7e0cb0a 100644 --- a/streamlit/interface_3.py +++ b/streamlit/interface_3.py @@ -20,7 +20,7 @@ import util import anthropic import re -REDIRECT_URI = 'http://172.29.20.126:8503' +REDIRECT_URI = 'https://doczy.aarete.com:8503' user_list = ['maamseek@aarete.com', 'smahdavian@aarete.com', 'ahinge@aarete.com', 'akadam@aarete.com' , 'piragavarapu@aarete.com', 'umistry@aarete.com', 'ahutchison@aarete.com', 'bgrunst@aarete.com', 'ddimeglio@aarete.com' , 'vnair@aarete.com'] @@ -39,9 +39,16 @@ with st.sidebar: 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: +try: + util.setup_page(REDIRECT_URI) + _,c1= st.columns([4,1]) + c1.write(f"User: **{st.session_state.user_info['displayName']}**") + user_mail = st.session_state.user_info['mail'] +except: + user_mail = 'maamseek@aarete.com' + + +if user_mail in user_list: field_row = st.columns([0.15, 0.45, 0.4]) with field_row[0]: