Team Ai
Modelpublic

CyberPeace-Institute/Cybersecurity-Knowledge-Graph

sourceHugging Facemitupdated 3y agoView on Hugging Face
23likes60downloads
event_arg_predict.py281 linesDownload Raw Back to root
1import streamlit as st2from annotated_text import annotated_text3import torch4from torch.utils.data import DataLoader5 6from .args_model_utils import tokenize_and_align_labels_with_pos_ner_dep, find_nearest_nugget_features, find_dep_depth7from .nugget_model_utils import CustomRobertaWithPOS8from .utils import get_content, get_event_nugget, get_idxs_from_text, get_entity_from_idx, list_of_pos_tags, event_args_list9 10from .event_nugget_predict import get_event_nuggets11import spacy12from transformers import AutoTokenizer13from datasets import load_dataset, Features, ClassLabel, Value, Sequence, Dataset14import os15 16os.environ["TOKENIZERS_PARALLELISM"] = "true"17 18def find_dep_depth(token):19    depth = 020    current_token = token21    while current_token.head != current_token:22        depth += 123        current_token = current_token.head24    return min(depth, 16)25 26 27nlp = spacy.load('en_core_web_sm')28 29pos_spacy_tag_list = ["ADJ","ADP","ADV","AUX","CCONJ","DET","INTJ","NOUN","NUM","PART","PRON","PROPN","PUNCT","SCONJ","SYM","VERB","SPACE","X"]30ner_spacy_tag_list = [bio + entity for entity in list(nlp.get_pipe('ner').labels) for bio in ["B-", "I-"]] + ["O"]31dep_spacy_tag_list = list(nlp.get_pipe("parser").labels)32event_nugget_tag_list = ["Databreach", "Ransom", "PatchVulnerability", "Phishing", "DiscoverVulnerability"]33arg_nugget_relative_pos_tag_list = ["before-same-sentence", "before-differ-sentence", "after-same-sentence", "after-differ-sentence"]34 35device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")36 37model_checkpoint = "ehsanaghaei/SecureBERT"38tokenizer = AutoTokenizer.from_pretrained(model_checkpoint, add_prefix_space=True)39 40# from .args_model_utils import CustomRobertaWithPOS as ArgumentModel41# model_nugget = ArgumentModel(num_classes=43)42# model_nugget.load_state_dict(torch.load(f"{os.path.dirname(os.path.abspath(__file__))}/argument_model_state_dict.pth", map_location=device)) 43# model_nugget.eval()44 45"""46Function: create_dataloader(text_input)47Description: This function creates a DataLoader for processing text data, tokenizes it, and organizes it into batches.48Inputs:49    - text_input: The input text to be processed.50Output:51    - dataloader: A DataLoader for the tokenized and batched text data.52    - tokenized_dataset_ner: The tokenized dataset used for training.53"""54def create_dataloader(model_nugget, text_input):55 56    event_nuggets = get_event_nuggets(model_nugget, text_input)57    doc = nlp(text_input)58 59    content_as_words_emdash = [tok.text for tok in doc]60    content_as_words_emdash = [word.replace("``", '"').replace("''", '"').replace("$", "") for word in content_as_words_emdash]61    content_idx_dict = get_idxs_from_text(text_input, content_as_words_emdash)62 63    data = []64 65    words = []66    arg_nugget_nearest_subtype = []67    arg_nugget_nearest_dist = []68    arg_nugget_relative_pos = []69 70    pos_spacy = [tok.pos_ for tok in doc]71    ner_spacy = [ent.ent_iob_ + "-" + ent.ent_type_ if ent.ent_iob_ != "O" else ent.ent_iob_ for ent in doc]72    dep_spacy = [tok.dep_ for tok in doc]73    depth_spacy = [find_dep_depth(tok) for tok in doc]74 75    for content_dict in content_idx_dict:76        start_idx, end_idx = content_dict["start_idx"], content_dict["end_idx"]77        nearest_subtype, nearest_dist, relative_pos = find_nearest_nugget_features(doc, content_dict["start_idx"], content_dict["end_idx"], event_nuggets)78        words.append(content_dict["word"])79 80        arg_nugget_nearest_subtype.append(nearest_subtype)81        arg_nugget_nearest_dist.append(nearest_dist)82        arg_nugget_relative_pos.append(relative_pos)83 84 85    content_token_len = len(tokenizer(words, truncation=False, is_split_into_words=True)["input_ids"])86    if content_token_len > tokenizer.model_max_length:87        no_split = (content_token_len // tokenizer.model_max_length) + 288        split_len = (len(words) // no_split) + 189 90        last_id = 091        threshold = split_len92 93        for id, token in enumerate(words):94            if token == "." and id > threshold:95                data.append(96                    {97                        "tokens" : words[last_id : id + 1],98                        "pos_spacy" : pos_spacy[last_id : id + 1],99                        "ner_spacy" : ner_spacy[last_id : id + 1],100                        "dep_spacy" : dep_spacy[last_id : id + 1],101                        "depth_spacy" : depth_spacy[last_id : id + 1],102                        "nearest_nugget_subtype" : arg_nugget_nearest_subtype[last_id : id + 1],103                        "nearest_nugget_dist" : arg_nugget_nearest_dist[last_id : id + 1],104                        "arg_nugget_relative_pos" : arg_nugget_relative_pos[last_id : id + 1]105                    }106                )107                last_id = id + 1108                threshold += split_len109        data.append({"tokens" : words[last_id : ],110                     "pos_spacy" : pos_spacy[last_id : ],111                     "ner_spacy" : ner_spacy[last_id : ],112                     "dep_spacy" : dep_spacy[last_id : ],113                     "depth_spacy" : depth_spacy[last_id : ],114                     "nearest_nugget_subtype" : arg_nugget_nearest_subtype[last_id : ],115                    "nearest_nugget_dist" : arg_nugget_nearest_dist[last_id : ],116                    "arg_nugget_relative_pos" : arg_nugget_relative_pos[last_id : ]}) 117    else:118        data.append(119            {120                "tokens" : words,121                "pos_spacy" : pos_spacy,122                "ner_spacy" : ner_spacy,123                "dep_spacy" : dep_spacy,124                "depth_spacy" : depth_spacy,125                "nearest_nugget_subtype" : arg_nugget_nearest_subtype,126                "nearest_nugget_dist" : arg_nugget_nearest_dist,127                "arg_nugget_relative_pos" : arg_nugget_relative_pos128            }129        )130 131 132    ner_features = Features({'tokens' : Sequence(feature=Value(dtype='string', id=None), length=-1, id=None),133                            'pos_spacy' : Sequence(feature=ClassLabel(num_classes=len(pos_spacy_tag_list), names=pos_spacy_tag_list, names_file=None, id=None), length=-1, id=None),134                            'ner_spacy' : Sequence(feature=ClassLabel(num_classes=len(ner_spacy_tag_list), names=ner_spacy_tag_list, names_file=None, id=None), length=-1, id=None),135                            'dep_spacy' : Sequence(feature=ClassLabel(num_classes=len(dep_spacy_tag_list), names=dep_spacy_tag_list, names_file=None, id=None), length=-1, id=None),136                            'depth_spacy' : Sequence(feature=ClassLabel(num_classes=17, names= list(range(17)), names_file=None, id=None), length=-1, id=None),137                            'nearest_nugget_subtype' : Sequence(feature=ClassLabel(num_classes=len(event_nugget_tag_list), names=event_nugget_tag_list, names_file=None, id=None), length=-1, id=None),138                            'nearest_nugget_dist' : Sequence(feature=ClassLabel(num_classes=11, names=list(range(11)), names_file=None, id=None), length=-1, id=None),139                            'arg_nugget_relative_pos' : Sequence(feature=ClassLabel(num_classes=len(arg_nugget_relative_pos_tag_list), names=arg_nugget_relative_pos_tag_list, names_file=None, id=None), length=-1, id=None),140                            })141 142    dataset = Dataset.from_list(data, features=ner_features)143    tokenized_dataset_ner = dataset.map(tokenize_and_align_labels_with_pos_ner_dep, fn_kwargs={'tokenizer' : tokenizer}, batched=True, load_from_cache_file=False)144    tokenized_dataset_ner = tokenized_dataset_ner.with_format("torch")145 146    tokenized_dataset_ner = tokenized_dataset_ner.remove_columns("tokens")147 148    batch_size = 4 # Number of input texts149    dataloader = DataLoader(tokenized_dataset_ner, batch_size=batch_size)150    return dataloader, tokenized_dataset_ner151 152"""153Function: predict(dataloader)154Description: This function performs prediction on a given dataloader using a trained model for label classification.155Inputs:156    - dataloader: A DataLoader containing the input data for prediction.157Output:158    - predicted_label: A tensor containing the predicted labels for each input in the dataloader.159"""160def predict(dataloader):161    predicted_label = []162    for batch in dataloader:163        with torch.no_grad():164            logits = model_nugget(**batch)165 166        batch_predicted_label = logits.argmax(-1)167        predicted_label.append(batch_predicted_label)168    return torch.cat(predicted_label, dim=-1)169 170"""171Function: show_annotations(text_input)172Description: This function displays annotated event arguments in the provided input text.173Inputs:174    - text_input: The input text containing event arguments to be annotated and displayed.175Output:176    - An interactive display of annotated event arguments within the input text.177"""178def show_annotations(text_input):179    st.title("Event Arguments")180 181    dataloader, tokenized_dataset_ner = create_dataloader(text_input)182    predicted_label = predict(dataloader)183 184    for idx, labels in enumerate(predicted_label):185        token_mask = [token > 2 for token in tokenized_dataset_ner[idx]["input_ids"]]186 187        tokens = tokenizer.convert_ids_to_tokens(tokenized_dataset_ner[idx]["input_ids"][token_mask], skip_special_tokens=True)188        tokens = [token.replace("Ġ", "").replace("Ċ", "").replace("âĢĻ", "'") for token in tokens]189 190        text = tokenizer.decode(tokenized_dataset_ner[idx]["input_ids"][token_mask])191        idxs = get_idxs_from_text(text, tokens)192 193        labels = labels[token_mask]194 195        annotated_text_list = []196        last_label = ""197        cumulative_tokens = "" 198        last_id = 0199 200        for idx, label in zip(idxs, labels):201            to_label = event_args_list[label]202            label_short = to_label.split("-")[1] if "-" in to_label else to_label203            if last_label == label_short:204                cumulative_tokens += text[last_id : idx["end_idx"]]205                last_id = idx["end_idx"]206            else:207                if last_label != "":208                    if last_label == "O":209                        annotated_text_list.append(cumulative_tokens)210                    else:211                        annotated_text_list.append((cumulative_tokens, last_label))212                last_label = label_short213                cumulative_tokens = idx["word"]214                last_id = idx["end_idx"]215        if last_label == "O":216            annotated_text_list.append(cumulative_tokens)217        else:  218            annotated_text_list.append((cumulative_tokens, last_label))219 220        annotated_text(annotated_text_list)221 222"""223Function: get_event_args(text_input)224Description: This function extracts predicted event arguments (event nuggets) from the provided input text.225Inputs:226    - text_input: The input text containing event nuggets to be extracted.227Output:228    - predicted_event_nuggets: A list of dictionaries, each representing an extracted event nugget with start and end offsets,229      subtype, and text content.230"""231def get_event_args(text_input):232    dataloader, tokenized_dataset_ner = create_dataloader(text_input)233    predicted_label = predict(dataloader)234 235    predicted_event_nuggets = []236    text_length = 0 237    for idx, labels in enumerate(predicted_label):238        token_mask = [token > 2 for token in tokenized_dataset_ner[idx]["input_ids"]]239 240        tokens = tokenizer.convert_ids_to_tokens(tokenized_dataset_ner[idx]["input_ids"][token_mask], skip_special_tokens=True)241        tokens = [token.replace("Ġ", "").replace("Ċ", "").replace("âĢĻ", "'") for token in tokens]242 243        text = tokenizer.decode(tokenized_dataset_ner[idx]["input_ids"][token_mask])244        idxs = get_idxs_from_text(text_input[text_length : ], tokens)245 246        labels = labels[token_mask]247 248        start_idx = 0249        end_idx = 0250        last_label = ""251 252        for idx, label in zip(idxs, labels):253            to_label = event_args_list[label]254            if "-" in to_label:255                label_split = to_label.split("-")[1]256            else:257                label_split = to_label258            259            if label_split == last_label:260                end_idx = idx["end_idx"]261            else:262                if text_input[start_idx : end_idx] != "" and last_label != "O":263                    predicted_event_nuggets.append(264                        {265                            "startOffset" : text_length + start_idx,266                            "endOffset" : text_length + end_idx,267                            "subtype" : last_label,268                            "text" : text_input[text_length + start_idx : text_length + end_idx]269                        }270                    )271                start_idx = idx["start_idx"]272                end_idx = idx["start_idx"] + len(idx["word"])273            last_label = label_split274        text_length += idx["end_idx"]275    return predicted_event_nuggets276   277 278 279 280 281