PRENT/PRENT-Codebook
0
1import json2import string3from time import time4 5import en_core_web_lg6import inflect7import nltk8import numpy as np9import pandas as pd10import streamlit as st11from nltk.tokenize import sent_tokenize12from transformers import pipeline13 14# Set constant values15INFLECT_ENGINE = inflect.engine()16TOP_K = 3017NLI_LIMIT = 0.918 19st.set_page_config(layout="wide")20 21 22def get_top_k():23 return TOP_K24 25 26def get_nli_limit():27 return NLI_LIMIT28 29 30### Streamlit specific31@st.cache(allow_output_mutation=True)32def load_model_prompting():33 return pipeline("fill-mask", model="distilbert-base-uncased")34 35 36@st.cache(allow_output_mutation=True)37def load_model_nli():38 try:39 return pipeline(40 task="sentiment-analysis", model="roberta-large-mnli", device="mps"41 )42 except:43 return pipeline(task="sentiment-analysis", model="roberta-large-mnli")44 45 46@st.cache(allow_output_mutation=True)47def load_spacy_pipeline():48 return en_core_web_lg.load()49 50 51@st.cache()52def download_punkt():53 nltk.download("punkt")54 55 56download_punkt()57 58 59@st.experimental_memo(max_entries=1)60def read_json_from_web(uploaded_json):61 return json.load(uploaded_json)62 63 64@st.experimental_memo(max_entries=1)65def read_csv_from_web(uploaded_file):66 """Read CSV from the streamlit interface67 68 :param uploaded_file: File to read69 :type uploaded_file: UploadedFile (BytesIO)70 :return: Dataframe71 :rtype: pandas DataFrame72 """73 try:74 # Try first to read comma separated and semicolon separated files75 data = pd.read_csv(uploaded_file, sep=None, engine="python")76 # If both are not correct, then it will error and go to the except77 except pd.errors.ParserError:78 # This should be the case when there is no separator (1 column csv)79 # Reset the IO object due to the previous crash80 uploaded_file.seek(0)81 # Use standard reading of CSV (no separator)82 data = pd.read_csv(uploaded_file)83 return data84 85 86def apply_style():87 # Avoid having ellipsis in the multi select options88 styl = """89 <style>90 .stMultiSelect span{91 max-width: none;92 93 }94 </style>95 """96 st.markdown(styl, unsafe_allow_html=True)97 98 # Set color of multiselect to red99 st.markdown(100 """101 <style>102 span[data-baseweb="tag"] {103 background-color: red !important;104 }105 </style>106 """,107 unsafe_allow_html=True,108 )109 110 hide_st_style = """111 <style>112 #MainMenu {visibility: hidden;}113 footer {visibility: hidden;}114 header {visibility: hidden;}115 </style>116 """117 st.markdown(hide_st_style, unsafe_allow_html=True)118 119 120def choose_text_menu(text):121 if "text" not in st.session_state:122 st.session_state.text = "Several demonstrators were injured."123 text = st.text_area("Event description", st.session_state.text)124 125 return text126 127 128def initiate_widget_st_state(widget_key, perm_key, default_value):129 if perm_key not in st.session_state:130 st.session_state[perm_key] = default_value131 if widget_key not in st.session_state:132 st.session_state[widget_key] = st.session_state[perm_key]133 134 135def get_idx_column(col_name, col_list):136 if col_name in col_list:137 return col_list.index(col_name)138 else:139 return 0140 141 142def callback_add_to_multiselect(str_to_add, multiselect_key, text_input_key, *keys):143 if len(str_to_add) == 0:144 st.warning("Word is empty, did you press Enter on the field text?")145 return146 current_dict = st.session_state147 *dict_keys, item_keys = keys148 try:149 for key in dict_keys:150 current_dict = current_dict[key]151 current_dict[item_keys].append(str_to_add)152 except KeyError as e:153 raise KeyError(keys) from e154 155 if multiselect_key in st.session_state:156 st.session_state[multiselect_key].append(str_to_add)157 else:158 st.session_state[multiselect_key] = [str_to_add]159 160 st.session_state[text_input_key] = ""161 162 163# Split the text into sentences. Necessary for NLI models164def split_sentences(text):165 return sent_tokenize(text)166 167 168def get_num_sentences_in_list_text(list_texts):169 num_sentences = 0170 for text in list_texts:171 num_sentences += len(split_sentences(text))172 return num_sentences173 174 175###### Prompting176def query_model_prompting(model, text, prompt_with_mask, top_k, targets):177 """Query the prompting model178 179 :param model: Prompting model object180 :type model: Huggingface pipeline object181 :param text: Event description (context)182 :type text: str183 :param prompt_with_mask: Prompt with a mask184 :type prompt_with_mask: str185 :param top_k: Number of tokens to output186 :type top_k: integer187 :param targets: Restrict the answer to these possible tokens188 :type targets: list189 :return: Results of the prompting model190 :rtype: list of dict191 """192 sequence = text + prompt_with_mask193 output_tokens = model(sequence, top_k=top_k, targets=targets)194 195 return output_tokens196 197 198def do_sentence_entailment(sentence, hypothesis, model):199 """Concatenate context and hypothesis then perform entailment200 201 :param sentence: Event description (context), 1 sentence202 :type sentence: str203 :param hypothesis: Mask filled with a token204 :type hypothesis: str205 :param model: NLI Model206 :type model: Huggingface pipeline207 :return: DataFrame containing the result of the entailment208 :rtype: pandas DataFrame209 """210 text = sentence + "</s></s>" + hypothesis211 res = model(text, return_all_scores=True)212 df_res = pd.DataFrame(res[0])213 df_res["label"] = df_res["label"].apply(lambda x: x.lower())214 df_res.columns = ["Label", "Score"]215 return df_res216 217 218def softmax(x):219 """Compute softmax values for each sets of scores in x."""220 return np.exp(x) / np.sum(np.exp(x), axis=0)221 222 223def get_singular_form(word):224 """Get the singular form of a word225 226 :param word: word227 :type word: string228 :return: singular form of the word229 :rtype: string230 """231 if INFLECT_ENGINE.singular_noun(word):232 return INFLECT_ENGINE.singular_noun(word)233 else:234 return word235 236 237######### NLI + PROMPTING238def do_text_entailment(text, hypothesis, model):239 """240 Do entailment for each sentence of the event description as241 model was trained on sentence pair242 243 :param text: Event Description (context)244 :type text: str245 :param hypothesis: Mask filled with a token246 :type hypothesis: str247 :param model: Model NLI248 :type model: Huggingface pipeline249 :return: List of entailment results for each sentence of the text250 :rtype: list251 """252 text_entailment_results = []253 for i, sentence in enumerate(split_sentences(text)):254 df_score = do_sentence_entailment(sentence, hypothesis, model)255 text_entailment_results.append((sentence, hypothesis, df_score))256 return text_entailment_results257 258 259def get_true_entailment(text_entailment_results, nli_limit):260 """261 From the result of each sentence entailment, extract the maximum entailment score and262 check if it's higher than the entailment threshold.263 """264 true_hypothesis_list = []265 max_score = 0266 for sentence_entailment in text_entailment_results:267 df_score = sentence_entailment[2]268 score = df_score[df_score["Label"] == "entailment"]["Score"].values.max()269 if score > max_score:270 max_score = score271 if max_score > nli_limit:272 true_hypothesis_list.append((sentence_entailment[1], np.round(max_score, 2)))273 return list(set(true_hypothesis_list))274 275 276def run_model_nli(data, batch_size, model_nli, use_tf=False):277 if not use_tf:278 return model_nli(data, top_k=3, batch_size=batch_size)279 else:280 raise NotImplementedError281 # return run_pipeline_on_gpu(data, batch_size, model_nli["tokenizer"], model_nli["model"])282 283 284def prompt_to_nli_batching(285 text,286 prompt,287 model_prompting,288 nli_model,289 nlp,290 top_k=10,291 nli_limit=0.5,292 targets=None,293 additional_words=None,294 remove_lemma=False,295 use_tf=False,296):297 # Check if text has end ponctuation298 if text[-1] not in string.punctuation:299 text += "."300 prompt_masked = prompt.format(model_prompting.tokenizer.mask_token)301 output_prompting = query_model_prompting(302 model_prompting, text, prompt_masked, top_k, targets=targets303 )304 if remove_lemma:305 output_prompting = filter_prompt_output_by_lemma(prompt, output_prompting, nlp)306 full_batch_concat = []307 prompt_tokens = []308 for token in output_prompting:309 hypothesis = prompt.format(token["token_str"])310 for i, sentence in enumerate(split_sentences(text)):311 full_batch_concat.append(sentence + "</s></s>" + hypothesis)312 prompt_tokens.append((token["token_str"], token["score"]))313 314 # Add words that must be tried for entailment315 # Also increase batch_size316 if additional_words:317 for i, sentence in enumerate(split_sentences(text)):318 for token in additional_words:319 hypothesis = prompt.format(token)320 full_batch_concat.append(sentence + "</s></s>" + hypothesis)321 prompt_tokens.append((token, 1))322 top_k = top_k + 1323 results_nli = run_model_nli(full_batch_concat, top_k, nli_model, use_tf)324 # Get entailed tokens325 entailed_tokens = []326 for i, res in enumerate(results_nli):327 entailed_tokens.extend(328 [329 (get_singular_form(prompt_tokens[i][0]), x["score"])330 for x in res331 if ((x["label"] == "ENTAILMENT") & (x["score"] > nli_limit))332 ]333 )334 if entailed_tokens:335 entailed_tokens = list(336 pd.DataFrame(entailed_tokens).groupby(0).max()[1].items()337 )338 339 return entailed_tokens, list(set(prompt_tokens))340 341 342def remove_similar_lemma_from_list(prompt, list_words, nlp):343 ## Compute a dictionnary with the lemma for all tokens344 ## If there is a duplicate lemma then the dictionnary value will be a list of the corresponding tokens345 lemma_dict = {}346 for each in list_words:347 mask_filled = nlp(prompt.strip(".").format(each))348 lemma_dict.setdefault([x.lemma_ for x in mask_filled][-1], []).append(each)349 350 ## Get back the list of tokens351 ## If multiple tokens available then take the shortest one352 new_token_list = []353 for key in lemma_dict.keys():354 if len(lemma_dict[key]) >= 1:355 new_token_list.append(min(lemma_dict[key], key=len))356 else:357 raise ValueError("Lemma dict has 0 corresponding words")358 return new_token_list359 360 361def filter_prompt_output_by_lemma(prompt, output_prompting, nlp):362 """363 Remove all similar lemmas from the prompt output (e.g. "protest", "protests")364 """365 list_words = [x["token_str"] for x in output_prompting]366 new_token_list = remove_similar_lemma_from_list(prompt, list_words, nlp)367 return [x for x in output_prompting if x["token_str"] in new_token_list]368 369 370# Streamlit specific run functions371@st.experimental_memo(max_entries=1024)372def do_prent(text, template, top_k, nli_limit, additional_words=None):373 """Function used to execute PRENT model374 375 :param text: Event text376 :type text: string377 :param template: Template with mask378 :type template: string379 :param top_k: Maximum tokens to output from prompting model380 :type top_k: int381 :param nli_limit: Threshold of entailment for NLI [0,1]382 :type nli_limit: float383 :param additional_words: List of words that bypass prompting and goes directly to NLI, defaults to None384 :type additional_words: list, optional385 :return: (Results Entailment, Results Prompting)386 :rtype: tuple387 """388 results_nli, results_pr = prompt_to_nli_batching(389 text,390 template,391 load_model_prompting(),392 load_model_nli(),393 load_spacy_pipeline(),394 top_k=top_k,395 nli_limit=nli_limit,396 targets=None,397 additional_words=additional_words,398 remove_lemma=True,399 )400 return results_nli, results_pr401 402 403def get_additional_words():404 """Extract the additional words from the codebook405 406 :return: list of additional words407 :rtype: list408 """409 if "add_words" in st.session_state.codebook:410 additional_words = st.session_state.codebook["add_words"]411 else:412 additional_words = None413 return additional_words414 415 416def run_prent(417 text="", templates=[], additional_words=None, progress=True, display_text=True418):419 """Execute PRENT over a list of templates and display streamlit widgets420 421 :param text: Event description, defaults to ""422 :type text: str, optional423 :param templates: Templates with a mask, defaults to []424 :type templates: list, optional425 :param additional_words: List of words to bypass prompting, defaults to None426 :type additional_words: list, optional427 :param progress: Display or not the progress bar, defaults to True428 :type progress: bool, optional429 :return: (results of prent, computation time)430 :rtype: tuple431 """432 # Check if there is any template and event description available433 if not templates:434 st.warning("Template list is empty. Please add one.")435 return None, None436 if not text:437 st.warning("Event description is empty.")438 return None, None439 440 # Display text only when computing441 if display_text:442 temp_text = st.empty()443 temp_text.markdown("**Event Descriptions:** {}".format(text))444 445 # Start progress bar446 if progress:447 progress_bar = st.progress(0)448 num_prent_call = len(templates)449 num_sentences = get_num_sentences_in_list_text([text])450 iter = 0451 t0 = time()452 453 # We set the radio choice of streamlit to Ignore at first454 if "accept_reject_text_perm" in st.session_state:455 st.session_state["accept_reject_text_perm"] = "Ignore"456 457 res = {}458 for template in templates:459 template = template.replace("[Z]", "{}")460 results_nli, results_pr = do_prent(461 text,462 template,463 top_k=TOP_K,464 nli_limit=NLI_LIMIT,465 additional_words=additional_words,466 )467 # Results_nli contains % of entailment, we only care about the tokens string468 res[template] = [x[0] for x in results_nli]469 470 # Update progress bar471 iter += 1472 if progress:473 progress_bar.progress((1 / num_prent_call) * (iter))474 if display_text:475 temp_text.markdown("")476 time_comput = (time() - t0) / num_sentences477 # This check is done otherwise the time of computation is replaced by the478 # time of computation when using cached value479 if not time_comput < st.session_state.time_comput / 5:480 st.session_state.time_comput = int(time_comput)481 482 # Store some results483 res["templates_used"] = templates484 res["additional_words_used"] = additional_words485 return res, time_comput486 487 488####### Find event types based on codebook and PRENT results489def check_any_conds(cond_any, list_res):490 """Function that evaluates the "OR" conditions of the codebook versus the list of filled templates491 492 :param cond_any: List of groundtruth filled templates493 :type cond_any: list494 :param list_res: A list of the filled templates given by PRENT495 :type list_res: list496 :return: True if any groundtruth template is inside the list given by PRENT497 :rtype: bool498 """499 cond_any = list(cond_any)500 condition = False501 # Return False if there is no any condition502 if not cond_any:503 return False504 for cond in cond_any:505 # With the current codebook design, this should never be true.506 # Before it was possible to have recursion to check AND conditions inside an OR condition507 if isinstance(cond, dict):508 condition = check_all_conds(cond["all"], list_res)509 else:510 # Check lowercase version of templates511 if cond.lower() in [x.lower() for x in list_res]:512 condition = True513 # Exit function as the other templates won't change the outcome514 return condition515 return condition516 517 518def check_all_conds(cond_all, list_res):519 """Function that evaluates the "AND" conditions of the codebook versus the list of filled templates520 521 :param cond_all: List of groundtruth filled templates522 :type cond_all: list523 :param list_res: A list of the filled templates given by PRENT524 :type list_res: list525 :return: True if all groundtruth template are inside the list given by PRENT526 :rtype: bool527 """528 cond_all = list(cond_all)529 # Return False if there is no all condition530 if not cond_all:531 return False532 # Start bool on True, and put it to false if any template is missing533 condition = True534 for cond in cond_all:535 # With the current codebook design, this should never be true.536 # Before it was possible to have recursion to check OR conditions inside an AND condition537 if isinstance(cond, dict):538 condition = check_any_conds(cond["any"])539 else:540 # Check lowercase version of templates541 if not (cond.lower() in [x.lower() for x in list_res]):542 condition = False543 # Exit function as the other templates won't change the outcome544 return condition545 return condition546 547 548def find_event_types(codebook, list_res):549 """This function evaluates the codebook and then outputs a list of events types corresponding to the given results of PRENT (list of filled templates).550 551 :param codebook: A codebook in the format given by the dashboard552 :type codebook: dict553 :param list_res: A list of the filled templates given by PRENT554 :type list_res: list555 :return: List of event type556 :rtype: list557 """558 list_event_type = []559 # Iterate over all defined event types560 for event_type in codebook["events"]:561 code_event = codebook["events"][event_type]562 563 is_not_all_event, is_not_any_event, is_not_event = False, False, False564 is_all_event, is_any_event, is_event = False, False, False565 566 # First check if NOT conditions are met567 # e.g. a filled template that is contrary to the event is present568 if "not_all" in code_event:569 cond_all = code_event["not_all"]570 if check_all_conds(cond_all, list_res):571 is_not_all_event = True572 if "not_any" in code_event:573 cond_any = code_event["not_any"]574 if check_any_conds(cond_any, list_res):575 is_not_any_event = True576 577 # Next we need to check if the "not_all" and "not_any" are related578 # by an "OR" or "AND".579 # This latest case needs special care because one of two list can580 # be empty so False581 if code_event["not_all_any_rel"] == "AND":582 if is_not_all_event and (not code_event["not_any"]):583 # If all TRUE and ANY is empty (so false)584 is_not_event = True585 elif is_not_any_event and (not code_event["not_all"]):586 # If any TRUE and ALL is empty (so false)587 is_not_event = True588 if is_not_all_event and is_not_any_event:589 is_not_event = True590 elif code_event["not_all_any_rel"] == "OR":591 if is_not_all_event or is_not_any_event:592 is_not_event = True593 594 # The other checks are not necessary if this is true, so we go595 # to the next iteration596 if is_not_event:597 continue598 599 # Similar to the previous checks but this time we look for templates that should be present600 if "all" in code_event:601 cond_all = code_event["all"]602 ## Then check if All conditions are met, if not exit603 if check_all_conds(cond_all, list_res):604 is_all_event = True605 if "any" in code_event:606 ## Finally check if Any conditions is met, if not exit607 cond_any = code_event["any"]608 if check_any_conds(cond_any, list_res):609 is_any_event = True610 611 # This case needs special care because one of two list can612 # be empty so False613 if code_event["all_any_rel"] == "AND":614 if is_all_event and (not code_event["any"]):615 # If all TRUE and ANY is empty (so false)616 is_event = True617 elif is_any_event and (not code_event["all"]):618 # If any TRUE and ALL is empty (so false)619 is_event = True620 elif is_all_event and is_any_event:621 is_event = True622 elif code_event["all_any_rel"] == "OR":623 if is_all_event or is_any_event:624 is_event = True625 626 # If all checks are correct, then we can add the event type to the output list627 if is_event:628 list_event_type.append(event_type)629 630 return list_event_type631 