Team Ai
Apppublic

clef/PRENT-Codebook

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
3_Apply_Codebook.py223 linesDownload Raw Back to pages
1import json2import os3import sys4 5import pandas as pd6import streamlit as st7 8current = os.path.dirname(os.path.realpath(__file__))9parent = os.path.dirname(current)10sys.path.append(parent)11from helpers import (12    apply_style,13    find_event_types,14    get_additional_words,15    get_nli_limit,16    get_num_sentences_in_list_text,17    get_top_k,18    run_prent,19)20 21### Styling22apply_style()23 24 25TOP_K = get_top_k()26NLI_LIMIT = get_nli_limit()27 28 29### Initialize session state variables30if "codebook" not in st.session_state:31    st.session_state.codebook = {}32    st.session_state.codebook.setdefault("events", {})33 34if "text" not in st.session_state:35    st.session_state.text = ""36 37if "res" not in st.session_state:38    st.session_state.res = None39 40if "accept_reject_text_perm" not in st.session_state:41    st.session_state.accept_reject_text_perm = None42 43if "validated_data" not in st.session_state:44    st.session_state["validated_data"] = {}45 46if "time_comput" not in st.session_state:47    st.session_state.time_comput = 2048 49if "rerun" not in st.session_state:50    st.session_state.rerun = False51 52if "label_res" not in st.session_state:53    st.session_state.label_res = {}54 55if "filtered_df" not in st.session_state:56    st.session_state["filtered_df"] = pd.DataFrame()57 58if len(st.session_state["filtered_df"]) == 0:59    st.warning("No data loaded.")60 61 62def reset_computation_results():63    st.session_state.res = {}64    st.session_state.recompute_all_templates = True65    st.session_state["accept_reject_text_perm"] = "Ignore"66    st.session_state.rerun = True67 68 69with st.sidebar:70    st.markdown(71        "Clicking any of these button during labeling will pause the process and download the latest version."72    )73    dl_labeled_button = st.empty()74    dl_labeled_button.download_button(75        label="Download Labeled Data",76        data=st.session_state["filtered_df"].to_csv(sep=";").encode("utf-8"),77        file_name="labeled_data.csv",78        mime="text/csv",79    )80 81    dl_prent_button = st.empty()82    dl_prent_button.download_button(83        label="Download PR-ENT results",84        data=json.dumps(st.session_state["label_res"], indent=3).encode("ASCII"),85        file_name="prent_results.json",86        mime="application/json",87    )88 89 90st.markdown(91    """# Apply codebook to the dataset92The currently loaded codebook will be used to find the event types of all event description in the currently loaded dataset. This can take some time (minutes to hours) depending on the size of the dataset (number of events, length of text).93 94 95"""96)97 98markdown_num_events = st.empty()99 100label_button = st.empty()101st.markdown("#### Main progress bar")102main_progress_bar = st.empty()103main_progress_bar = main_progress_bar.progress(0)104 105st.markdown("#### Last labeled event")106temp_text = st.empty()107temp_class = st.empty()108temp_text.markdown("**Event Descriptions:** {}".format(""))109temp_class.markdown("**Event Types Classification**: {}".format(""))110st.markdown(111    """#### Pause/Stop the event coding112Pressing the button once will stop the process at the next iteration."""113)114stop_button = st.button("Stop")115 116for event_type in st.session_state.codebook["events"]:117    if event_type not in st.session_state.filtered_df.columns:118        st.session_state.filtered_df[event_type] = 0119 120expected_time = 0121num_sentences = 0122for idx in st.session_state.filtered_df.index:123    subsampled_data = st.session_state.filtered_df.loc[idx:idx]124    list_text = subsampled_data[st.session_state["text_column_design_perm"]].values[:1]125    list_index = subsampled_data.index[:1]126    if list_text[0] != st.session_state.text:127        reset_computation_results()128    st.session_state.text = list_text[0]129    num_sentences += get_num_sentences_in_list_text([st.session_state.text])130    expected_time += st.session_state.time_comput * get_num_sentences_in_list_text(131        [st.session_state.text]132    )133 134markdown_num_events.markdown(135    "Number of events: {} ¦ Number of sentences: {}".format(136        len(st.session_state.filtered_df.index), num_sentences137    )138)139 140 141if label_button.button(142    "Label Data", disabled=len(st.session_state["filtered_df"]) == 0143):144    num_text = 0145    main_progress_bar.progress(num_text)146    temp_text.markdown("")147    temp_class.markdown("")148    tot_num_text = len(st.session_state.filtered_df.index)149 150    for idx in st.session_state.filtered_df.index:151        subsampled_data = st.session_state.filtered_df.loc[idx:idx]152        list_text = subsampled_data[st.session_state["text_column_design_perm"]].values[153            :1154        ]155        list_index = subsampled_data.index[:1]156        if list_text[0] != st.session_state.text:157            reset_computation_results()158        st.session_state.text = list_text[0]159        st.session_state.text_idx = list_index[0]160        st.session_state.template_list = []161        st.session_state.text_display = st.session_state.text162 163        st.session_state.res = {}164        res, time_comput = run_prent(165            st.session_state.text,166            st.session_state.codebook["templates"],167            get_additional_words(),168            progress=False,169            display_text=False,170        )171        st.session_state.res = res172 173        list_filled_templates = []174        for template in st.session_state.res:175            tmp = template.replace("[Z]", "{}")176            list_filled_templates.extend(177                [tmp.format(x) for x in st.session_state.res[template]]178            )179        list_event_type = find_event_types(180            st.session_state.codebook, list_filled_templates181        )182        for event_type in list_event_type:183            st.session_state.filtered_df.loc[idx, event_type] = 1184        temp_text.markdown(185            "**Event Descriptions:** {}".format(st.session_state.text_display)186        )187        temp_class.markdown(188            "**Event Types Classification**: {}".format("; ".join(list_event_type))189        )190 191        # Save results192        st.session_state.label_res[st.session_state.text_display] = {}193        st.session_state.label_res[st.session_state.text_display][194            "prent_results"195        ] = st.session_state.res196        st.session_state.label_res[st.session_state.text_display]["prent_params"] = (197            TOP_K,198            NLI_LIMIT,199        )200        st.session_state.label_res[st.session_state.text_display][201            "event_types"202        ] = list_event_type203 204        num_text += 1205        main_progress_bar.progress(num_text / tot_num_text)206 207    # Need to update the buttons otherwise it doesn't update the downloaded file208    # and the user would need to click two times209    dl_labeled_button.download_button(210        label="Download Labeled Data",211        data=st.session_state["filtered_df"].to_csv(sep=";").encode("utf-8"),212        file_name="labeled_data.csv",213        mime="text/csv",214        key="tmp",215    )216 217    dl_prent_button.download_button(218        label="Download PR-ENT results",219        data=json.dumps(st.session_state["label_res"], indent=3).encode("ASCII"),220        file_name="prent_results.json",221        mime="application/json",222    )223