CyberPeace-Institute/Cybersecurity-Knowledge-Graph
2360
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 