clef/PRENT-Codebook
0
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 