clef/PRENT-Codebook
0
1import datetime as datetime2import hashlib3import json4import os5import sys6 7import pandas as pd8import streamlit as st9 10current = os.path.dirname(os.path.realpath(__file__))11parent = os.path.dirname(current)12sys.path.append(parent)13from helpers import (14 apply_style,15 callback_add_to_multiselect,16 choose_text_menu,17 do_prent,18 find_event_types,19 get_additional_words,20 get_idx_column,21 get_nli_limit,22 get_num_sentences_in_list_text,23 get_top_k,24 initiate_widget_st_state,25 run_prent,26)27 28# Set constant values29TOP_K = get_top_k()30NLI_LIMIT = get_nli_limit()31 32### Styling33# Needs to be done first34apply_style()35 36# Avoid having ellipsis in the multi select options37styl = """38 <style>39 .stMultiSelect span{40 max-width: none;41 42 }43 </style>44 """45st.markdown(styl, unsafe_allow_html=True)46 47# Set color of multiselect to red48st.markdown(49 """50 <style>51 span[data-baseweb="tag"] {52 background-color: red !important;53 }54 </style>55 """,56 unsafe_allow_html=True,57)58 59 60def validated_metric_per_event_types(validated_dataset):61 """Compute the accuracy metrics of the validated dataset62 for each event type. Compute True Positive, False Negative,63 True Negative, False Positive.64 65 :param validated_dataset: Dictionary containing results of PRENT validated by the user66 :type validated_dataset: dict67 :return: Dictionnary containing accuracy metric for all event types68 :rtype: dict69 """70 dict_acc = {}71 72 for key, val in validated_dataset.items():73 # Compute the event types based on the computed templates of PRENT74 pred_event_types = find_event_types(75 st.session_state.codebook, val["filled_templates"]76 )77 true_event_types = val["event_types"]78 # Compute only accuracy for accepted samples79 if val["decision"] == "Accept":80 # Iterate over all possible event types81 for event_type in st.session_state.codebook["events"].keys():82 dict_acc.setdefault(event_type, {})83 dict_acc[event_type].setdefault("TP", 0)84 dict_acc[event_type].setdefault("FN", 0)85 dict_acc[event_type].setdefault("FP", 0)86 dict_acc[event_type].setdefault("TN", 0)87 if (event_type in true_event_types) and (88 event_type in pred_event_types89 ):90 dict_acc[event_type]["TP"] += 191 elif (event_type in true_event_types) and not (92 event_type in pred_event_types93 ):94 dict_acc[event_type]["FN"] += 195 elif not (event_type in true_event_types) and (96 event_type in pred_event_types97 ):98 dict_acc[event_type]["FP"] += 199 else:100 dict_acc[event_type]["TN"] += 1101 102 # Normalize metrics103 if dict_acc:104 for event_type in st.session_state.codebook["events"].keys():105 dict_acc[event_type]["Accuracy"] = (106 dict_acc[event_type]["TP"] + dict_acc[event_type]["TN"]107 ) / (108 dict_acc[event_type]["TP"]109 + dict_acc[event_type]["TN"]110 + dict_acc[event_type]["FP"]111 + dict_acc[event_type]["FN"]112 )113 114 return dict_acc115 116 117def store_validated_data(118 text,119 decision,120 text_idx,121 templates,122 additional_words,123 list_event_type,124 prent_params=(TOP_K, NLI_LIMIT),125):126 """Function used to store the results of PRENT in a DataFrame and in the127 session state of Streamlit.128 129 :param text: Event description130 :type text: string131 :param decision: Decision of the user (Accept/Reject/Ignore)132 :type decision: string133 :param text_idx: Index of the event134 :type text_idx: int135 :param templates: List of template used136 :type templates: list137 :param additional_words: List of additional words used138 :type additional_words: list139 :param list_event_type: List of event type found by PRENT and Codebook140 :type list_event_type: list141 :param prent_params: Parameters of PRENT, defaults to (TOP_K, NLI_LIMIT)142 :type prent_params: tuple, optional143 """144 if "validated_data" not in st.session_state:145 st.session_state["validated_data"] = {}146 147 # Generate an index if the text is not coming from a csv148 if not text_idx:149 # Create a hash of 8 digits of the text to put as index150 data_idx = str(151 "manual_{}".format(152 int(153 hashlib.sha256(text.encode("utf-8")).hexdigest(),154 16,155 )156 % 10**8157 )158 )159 else:160 data_idx = str(text_idx)161 162 if data_idx not in st.session_state["validated_data"]:163 st.session_state["validated_data"][data_idx] = {}164 st.session_state["validated_data"][data_idx]["text"] = text165 st.session_state["validated_data"][data_idx]["templates"] = [166 template.replace("{}", "[Z]") for template in templates167 ]168 st.session_state["validated_data"][data_idx]["additional_words"] = additional_words169 st.session_state["validated_data"][data_idx]["event_types"] = list_event_type170 st.session_state["validated_data"][data_idx][171 "filled_templates"172 ] = list_filled_templates173 st.session_state["validated_data"][data_idx]["decision"] = decision174 st.session_state["validated_data"][data_idx]["prent_params"] = prent_params175 176 177### Initialize session state variables178if "codebook" not in st.session_state:179 st.session_state.codebook = {}180 st.session_state.codebook.setdefault("events", {})181 st.session_state.codebook["templates"] = []182if "text" not in st.session_state:183 st.session_state.text = ""184if "res" not in st.session_state:185 st.session_state.res = None186if "accept_reject_text_perm" not in st.session_state:187 st.session_state.accept_reject_text_perm = None188if "validated_data" not in st.session_state:189 st.session_state["validated_data"] = {}190if "time_comput" not in st.session_state:191 st.session_state.time_comput = 20192if "rerun" not in st.session_state:193 st.session_state.rerun = False194if "recompute_all_templates" not in st.session_state:195 st.session_state.recompute_all_templates = False196 197 198def reset_computation_results():199 """Reset cached values in session state related to computations"""200 st.session_state.res = {}201 st.session_state.recompute_all_templates = True202 st.session_state["accept_reject_text_perm"] = "Ignore"203 st.session_state.rerun = True204 205 206def get_all_filled_templates(results):207 """Create the filled templates from PRENT results. Merging template with mask208 with the entailed tokens.209 210 :param results: Dictionary containing PRENT results211 :type results: dict212 :return: List of all entailed templates213 :rtype: list214 """215 filled_templates = []216 templates_used = [x.replace("[Z]", "{}") for x in results["templates_used"]]217 for template in templates_used:218 filled_template = [template.format(x) for x in results[template]]219 filled_templates.extend(filled_template)220 221 return filled_templates222 223 224# Split streamlit dashboard225col_intro_left, col_intro_righter = st.columns([8, 8])226with col_intro_left:227 st.markdown(228 """ # Codebook Design229 """230 )231 232 233def load_demo(234 codebook_path="codebook_demo.json",235 validated_data_path="validated_data_demo.json",236 csv_data_path="data_demo.csv",237):238 """Load demonstration files from disk239 240 :param codebook_path: path to codebook, defaults to "codebook_demo.json"241 :type codebook_path: str, optional242 :param validated_data_path: path to validated dataset, defaults to "validated_data_demo.json"243 :type validated_data_path: str, optional244 :param csv_data_path: path to raw data, defaults to "data_demo.csv"245 :type csv_data_path: str, optional246 """247 st.session_state.codebook = json.load(open(codebook_path))248 st.session_state.validated_data = json.load(open(validated_data_path))249 st.session_state.data = pd.read_csv(csv_data_path, delimiter=";")250 st.session_state.filtered_df = st.session_state.data251 st.session_state.text_column_design_perm = "Event Descriptions"252 st.session_state["multiselect_classes"] = list(253 st.session_state.codebook["events"].keys()254 )255 st.session_state.text_idx = 0256 st.session_state.text = (257 "On 23 August, a group attacked a village, abducting 6 people."258 )259 st.session_state.text_display = (260 "On 23 August, a group attacked a village, abducting 6 people."261 )262 st.session_state["text_options_valid_perm"] = "From CSV"263 st.session_state["text_options_valid"] = "From CSV"264 265 266def clear_all():267 """Cleare session state"""268 for each in st.session_state:269 del st.session_state[each]270 st.experimental_rerun()271 272 273# Add two buttons in the sidebar to load and clear the demo274with st.sidebar:275 if st.button("Load Demo"):276 load_demo()277 278 if st.button("Clear Demo"):279 clear_all()280 281 st.write("********")282 283 284with st.sidebar:285 # Next function used for callback when download286 def update_codebook_save_time():287 st.session_state.save_codebook_time = (288 datetime.datetime.now().astimezone().strftime("%Y-%m-%d %H:%M:%S %z")289 )290 291 if st.download_button(292 label="Download codebook as JSON",293 data=json.dumps(st.session_state.codebook, indent=3).encode("ASCII"),294 file_name="codebook.json",295 mime="application/json",296 ):297 update_codebook_save_time()298 if "save_codebook_time" in st.session_state:299 st.write("Saved on: " + st.session_state.save_codebook_time)300 301 302with st.sidebar:303 # Next function used for callback when download304 def update_validated_save_time():305 st.session_state.save_validated_time = (306 datetime.datetime.now().astimezone().strftime("%Y-%m-%d %H:%M:%S %z")307 )308 309 if st.download_button(310 label="Download labeled data",311 data=json.dumps(st.session_state["validated_data"], indent=3).encode("ASCII"),312 file_name="validated_data.json",313 mime="application/json",314 ):315 update_validated_save_time()316 if "save_validated_time" in st.session_state:317 st.write("Saved on: " + st.session_state.save_validated_time)318 319# Add text to sidebar320with st.sidebar:321 st.write("********")322 st.markdown(323 """324#### Manual:325 3261. Set the list of possible event types3272. Select the input mode of the data (Manual or CSV)3283. If the codebook is empty, write a default template329 - `This event involves [Z].` is a good starting point3304. Write/Select an event description3315. Run PR-ENT3326. Check the event type classification333 - If it is correct then select Accept and return to step 4.334 - If it is wrong then select Reject and populate the codebook with the appropriate filled templates. The classification is updated for each change, when it is correct, click Accept.3357. Return to step 4336 337#### Tips & Tricks:338 339- If you start a codebook from scratch, it may be easier to pass a manual text example for each event type to get a first codebook draft340- Current codebook accuracy based on labeled data can be found in the top right341- The approach does not aim for perfect accuracy and some failures can happen, e.g. some event descriptions can produce filled templates that are not satisfactory.342 """343 )344 345# Add accuracy table346with col_intro_righter:347 accuracy = st.empty()348 # We fill the table with the last acc to avoid having it disappearing each time349 if "acc_df" in st.session_state:350 accuracy.table(351 st.session_state.acc_df.loc["Accuracy":"Accuracy"].style.format("{:.2}")352 )353 performance_container = st.expander("Detailed Performances")354 355 356st.write("*********")357col_left, col_right = st.columns(2)358 359# Add widgets to add event type and choose text input360with col_intro_left:361 with st.expander("Event Types List"):362 st.markdown(363 """364 ## Select Event Types.365 """366 )367 368 if "class_list_perm" not in st.session_state:369 st.session_state["class_list_perm"] = []370 371 # Text field + button to add new event types to multiselect372 new_class = st.text_input(373 "Add a new event type", "", key="new_class_text_input"374 )375 st.button(376 "Add Class",377 on_click=callback_add_to_multiselect,378 args=(379 new_class,380 "multiselect_classes",381 "new_class_text_input",382 "class_list_perm",383 ),384 )385 # Multiselect to choose event types386 if "multiselect_classes" not in st.session_state:387 st.session_state["multiselect_classes"] = list(388 st.session_state.codebook["events"].keys()389 )390 class_list = st.multiselect(391 "Event Type List",392 set(393 st.session_state["class_list_perm"]394 + list(st.session_state.codebook["events"].keys())395 ),396 st.session_state["multiselect_classes"],397 key="multiselect_classes",398 )399 st.session_state["class_list_perm"] = class_list400 401 with st.expander("Select Text Input Mode (Manual, CSV)"):402 st.write(403 """404 Choose the text input of the event descriptions. Three choices:405 - Manual: One event description can be manually input406 - From CSV: If a CSV of event descriptions was provided407 """408 )409 410 def callback_radio_text_choice():411 st.session_state.text = ""412 st.session_state.text_display = ""413 414 initiate_widget_st_state(415 "text_options_valid", "text_options_valid_perm", "Manual"416 )417 st.session_state["text_options_valid_perm"] = st.radio(418 "Choose text input",419 ["Manual", "From CSV"],420 index=get_idx_column(421 st.session_state["text_options_valid"], ["Manual", "From CSV"]422 ),423 key="text_options_valid",424 on_change=callback_radio_text_choice,425 horizontal=True,426 )427 428 429with col_left:430 if st.session_state["text_options_valid_perm"] == "Manual":431 text = choose_text_menu("")432 # Reset all computations if text has changed433 if text != st.session_state.text:434 reset_computation_results()435 st.session_state.text_idx = None436 st.session_state.text = text437 st.session_state.text_display = text438 elif st.session_state["text_options_valid_perm"] == "From CSV":439 if st.button("Select Random Text"):440 sample = st.session_state.filtered_df.sample(n=1).iloc[0]441 text = sample[st.session_state["text_column_design_perm"]]442 idx = sample.name443 if text != st.session_state.text:444 reset_computation_results()445 st.session_state.text = text446 st.session_state.text_idx = idx447 st.session_state.text_display = st.session_state.text448 449 expected_time = st.session_state.time_comput * get_num_sentences_in_list_text(450 [st.session_state.text]451 )452 if st.button("Run PR-ENT / Expected time: {}sec".format(expected_time)):453 if "templates" in st.session_state.codebook:454 templates = st.session_state.codebook["templates"]455 else:456 templates = []457 st.warning("No template in codebook. Please add one.")458 459 additional_words = get_additional_words()460 st.session_state.res = {}461 res, time_comput = run_prent(st.session_state.text, templates, additional_words)462 st.session_state.res = res463 464 st.write("**Event Descriptions:** {}".format(st.session_state.text_display))465 ev_desc = st.empty()466 radio_empty = st.empty()467 468 if st.session_state.res:469 list_filled_templates = get_all_filled_templates(st.session_state.res)470 471 list_event_type = find_event_types(472 st.session_state.codebook, list_filled_templates473 )474 event_type_text = ev_desc.markdown(475 "**Current Event Types Classification**: {}".format(476 "; ".join(list_event_type)477 )478 )479 480 if "accept_reject_text_perm" not in st.session_state:481 st.session_state["accept_reject_text_perm"] = "Ignore"482 483 def callback_function(mod, key):484 st.session_state[mod] = st.session_state[key]485 486 radio_empty.radio(487 "Accept or Reject Coding",488 ["Ignore", "Accept", "Reject"],489 key="accept_reject_text",490 on_change=callback_function,491 args=(492 "accept_reject_text_perm",493 "accept_reject_text",494 ),495 index=get_idx_column(496 st.session_state["accept_reject_text_perm"],497 ["Ignore", "Accept", "Reject"],498 ),499 horizontal=True,500 )501 502 decision = st.session_state["accept_reject_text_perm"]503 text_idx = st.session_state.text_idx504 text = st.session_state.text505 store_validated_data(506 text,507 decision,508 text_idx,509 st.session_state.res["templates_used"],510 st.session_state.res["additional_words_used"],511 list_event_type,512 prent_params=(TOP_K, NLI_LIMIT),513 )514 515 516with col_right:517 518 if (519 st.session_state["accept_reject_text_perm"] == "Reject"520 ) or not st.session_state.codebook["templates"]:521 with st.expander("Add Templates + Explanation"):522 st.markdown(523 """524 ## Add Templates525 """526 )527 st.markdown(528 """529 For each template added. PR-ENT will be run on the selected text.530 """531 )532 533 if "templates" not in st.session_state.codebook:534 st.session_state.codebook["templates"] = []535 536 template = st.text_input(537 "Template with a mask [Z].", "This event involves [Z]."538 )539 540 if st.button("Add template"):541 if template not in st.session_state.codebook["templates"]:542 ## Add template to codebook543 st.session_state.codebook["templates"].append(template)544 545 additional_words = get_additional_words()546 prompt = template.replace("[Z]", "{}")547 results_nli, _ = do_prent(548 st.session_state.text,549 prompt,550 TOP_K,551 NLI_LIMIT,552 additional_words,553 )554 tokens_nli = [x[0] for x in results_nli]555 556 # Update result table with new template557 if not st.session_state["res"]:558 st.session_state.res = {}559 st.session_state.res["additional_words_used"] = additional_words560 st.session_state.res["templates_used"] = []561 st.session_state.res[prompt] = tokens_nli562 st.session_state.res["templates_used"].append(template)563 st.write("Template '{}' added.".format(template))564 else:565 st.write("Template '{}' already added.".format(template))566 567 if st.session_state.codebook["templates"]:568 with st.expander("Populate Codebook Explanation"):569 st.markdown(570 """571 ## Set the filled template to each class.572 For each class you can select one or more filled templates. When the evaluation will573 be made, these templates will be compared with the results of PR-ENT. There are 4 options:574 - ALL: If **ALL** of these filled templates are present in the results of PR-ENT then this event type is correct575 - ANY: If **ANY** of these filled templates is present in the results of PR-ENT then this event type is correct576 - NOT ALL: If **ALL** of these filled templates are present in the results of PR-ENT, then this event type is **not** correct577 - e.g. You may want to remove all *explosions* events from a class *Killings*.578 - NOT ANY: If **ANY** of these filled templates is present in the results of PR-ENT, then this event type is **not** correct579 580 Moreover, **ANY/ALL** and **NOT ANY/ NOT ALL** can be made in relation by a **AND / OR** condition.581 """582 )583 584 st.write("***************")585 st.write("### Populate Codebook")586 if not class_list:587 st.warning("No event type in codebook.")588 589 tokens_list = get_all_filled_templates(st.session_state.res)590 591 for event_type in class_list:592 st.session_state.codebook["events"].setdefault(event_type, {})593 event_type_chosen = event_type594 with st.expander(event_type):595 596 def declare_ms_event_templates(597 widget_key, widget_display, codebook_key598 ):599 if widget_key not in st.session_state:600 st.session_state[widget_key] = st.session_state.codebook[601 "events"602 ][event_type_chosen].setdefault(codebook_key, [])603 604 tokens_all = st.multiselect(605 widget_display,606 set(607 list(608 tokens_list609 + st.session_state.codebook["events"][610 event_type_chosen611 ].setdefault(codebook_key, [])612 )613 ),614 st.session_state[widget_key],615 key=widget_key,616 )617 st.session_state.codebook["events"][event_type_chosen][618 codebook_key619 ] = tokens_all620 621 declare_ms_event_templates(622 "ms_all_{}".format(event_type_chosen), "ALL", "all"623 )624 625 st.session_state.codebook["events"][event_type_chosen][626 "all_any_rel"627 ] = st.selectbox(628 "Relation",629 ["AND", "OR"],630 index=get_idx_column(631 st.session_state.codebook["events"][632 event_type_chosen633 ].setdefault("all_any_rel", "OR"),634 ["AND", "OR"],635 ),636 key="select_relation_any_all_{}".format(event_type_chosen),637 )638 639 declare_ms_event_templates(640 "ms_any_{}".format(event_type_chosen), "ANY", "any"641 )642 643 declare_ms_event_templates(644 "ms_not_all_{}".format(event_type_chosen), "NOT ALL", "not_all"645 )646 647 st.session_state.codebook["events"][event_type_chosen][648 "not_all_any_rel"649 ] = st.selectbox(650 "Relation",651 ["AND", "OR"],652 index=get_idx_column(653 st.session_state.codebook["events"][654 event_type_chosen655 ].setdefault("not_all_any_rel", "OR"),656 ["AND", "OR"],657 ),658 key="select_relation_not_any_all_{}".format(event_type_chosen),659 )660 661 declare_ms_event_templates(662 "ms_not_any_{}".format(event_type_chosen), "NOT ANY", "not_any"663 )664 665 # Workaround to avoid the expanders closing after first modification666 # I have no explanation for the bug667 if st.session_state.rerun:668 st.session_state.rerun = False669 st.experimental_rerun()670 671 672if "validated_data" in st.session_state:673 recompute = False674 performance_container.markdown(675 "If a new template is added, the previous labeled samples needs to be recomputed with it. The next button allows that, however it can take some time depending on the number of samples."676 )677 if performance_container.button(678 "Recompute Missing Templates", key="recompute_temp"679 ):680 prog_bar = performance_container.progress(0)681 for i, datapoint in enumerate(st.session_state["validated_data"].values()):682 if not set(st.session_state.codebook["templates"]).issubset(683 set(datapoint["templates"])684 ):685 # Get templates that are missing from results but present in codebook686 # These happens if templates are added a posteriori687 missing_templates = list(688 set(st.session_state.codebook["templates"])689 - set(set(datapoint["templates"]))690 )691 recompute = True692 # For now additional words are not recomputed693 if not set(st.session_state.codebook["add_words"]).issubset(694 set(datapoint["additional_words"])695 ):696 missing_add_words = list(697 set(st.session_state.codebook["add_words"])698 - set(set(datapoint["additional_words"]))699 )700 recompute = True701 else:702 missing_add_words = None703 704 if recompute:705 res, _ = run_prent(706 datapoint["text"],707 missing_templates,708 missing_add_words,709 progress=False,710 )711 datapoint["filled_templates"].extend(get_all_filled_templates(res))712 datapoint["templates"].extend(missing_templates)713 prog_bar.progress(714 (1 / len(st.session_state["validated_data"].values())) * (i + 1)715 )716 717 st.session_state.acc_df = pd.DataFrame(718 validated_metric_per_event_types(st.session_state["validated_data"])719 )720 accuracy.table(721 st.session_state.acc_df.loc["Accuracy":"Accuracy"].style.format("{:.2}")722 )723 performance_container.markdown("### Performances on labeled dataset")724 performance_container.dataframe(st.session_state.acc_df.style.format("{:.3}"))725 726if st.session_state.res:727 list_filled_templates = get_all_filled_templates(st.session_state.res)728 list_event_type = find_event_types(st.session_state.codebook, list_filled_templates)729 ev_desc.markdown(730 "**Current Event Types Classification**: {}".format("; ".join(list_event_type))731 )732 