Team Ai
Apppublic

clef/PRENT-Codebook

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
1_Codebook_Design.py732 linesDownload Raw Back to pages
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