Team Ai
Apppublic

PRENT/PRENT-Codebook

sourceHugging Faceupdated 4y agoView on Hugging Face
0likes
helpers.py631 linesDownload Raw Back to root
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