CyberPeace-Institute/Cybersecurity-Knowledge-Graph
2360
1import torch2import spacy3import en_core_web_sm4from torch import nn5import math6 7 8device = torch.device("cuda") if torch.cuda.is_available() else torch.device("cpu")9 10from transformers import AutoModel, TrainingArguments, Trainer, RobertaTokenizer, RobertaModel11from transformers import AutoTokenizer12 13model_checkpoint = "ehsanaghaei/SecureBERT"14 15tokenizer = AutoTokenizer.from_pretrained(model_checkpoint, add_prefix_space=True)16roberta_model = RobertaModel.from_pretrained(model_checkpoint).to(device)17 18nlp = en_core_web_sm.load()19pos_spacy_tag_list = ["ADJ","ADP","ADV","AUX","CCONJ","DET","INTJ","NOUN","NUM","PART","PRON","PROPN","PUNCT","SCONJ","SYM","VERB","SPACE","X"]20ner_spacy_tag_list = [bio + entity for entity in list(nlp.get_pipe('ner').labels) for bio in ["B-", "I-"]] + ["O"]21 22 23class CustomRobertaWithPOS(nn.Module):24 def __init__(self, num_classes):25 super(CustomRobertaWithPOS, self).__init__()26 self.num_classes = num_classes27 self.pos_embed = nn.Embedding(len(pos_spacy_tag_list), 16)28 self.ner_embed = nn.Embedding(len(ner_spacy_tag_list), 16)29 self.roberta = roberta_model30 self.dropout1 = nn.Dropout(0.2)31 self.fc1 = nn.Linear(self.roberta.config.hidden_size, num_classes)32 33 def forward(self, input_ids, attention_mask, pos_spacy, ner_spacy, dep_spacy, depth_spacy):34 outputs = self.roberta(input_ids=input_ids, attention_mask=attention_mask)35 last_hidden_output = outputs.last_hidden_state36 37 pos_mask = pos_spacy != -10038 39 pos_one_hot = torch.zeros((pos_spacy.shape[0], pos_spacy.shape[1], len(pos_spacy_tag_list)), dtype=torch.long)40 pos_one_hot[pos_mask, pos_spacy[pos_mask]] = 141 pos_one_hot = pos_one_hot.to(device)42 43 ner_mask = ner_spacy != -10044 45 ner_one_hot = torch.zeros((ner_spacy.shape[0], ner_spacy.shape[1], len(ner_spacy_tag_list)), dtype=torch.long)46 ner_one_hot[ner_mask, ner_spacy[ner_mask]] = 147 ner_one_hot = ner_one_hot.to(device)48 49 features_concat = last_hidden_output50 features_concat = self.dropout1(features_concat)51 52 logits = self.fc1(features_concat)53 54 return logits55 56 57def tokenize_and_align_labels_with_pos_ner_dep(examples, tokenizer, label_all_tokens = True):58 tokenized_inputs = tokenizer(examples["tokens"], padding='max_length', truncation=True, is_split_into_words=True)59 #tokenized_inputs.pop('input_ids')60 ner_spacy = []61 pos_spacy = []62 dep_spacy = []63 depth_spacy = []64 65 for i, (pos, ner, dep, depth) in enumerate(zip(examples["pos_spacy"], 66 examples["ner_spacy"], 67 examples["dep_spacy"], 68 examples["depth_spacy"])):69 word_ids = tokenized_inputs.word_ids(batch_index=i)70 previous_word_idx = None71 ner_spacy_ids = []72 pos_spacy_ids = []73 dep_spacy_ids = []74 depth_spacy_ids = []75 76 for word_idx in word_ids:77 # Special tokens have a word id that is None. We set the label to -100 so they are automatically78 # ignored in the loss function.79 if word_idx is None:80 ner_spacy_ids.append(-100)81 pos_spacy_ids.append(-100)82 dep_spacy_ids.append(-100)83 depth_spacy_ids.append(-100)84 # We set the label for the first token of each word.85 elif word_idx != previous_word_idx:86 ner_spacy_ids.append(ner[word_idx])87 pos_spacy_ids.append(pos[word_idx])88 dep_spacy_ids.append(dep[word_idx])89 depth_spacy_ids.append(depth[word_idx])90 # For the other tokens in a word, we set the label to either the current label or -100, depending on91 # the label_all_tokens flag.92 else:93 ner_spacy_ids.append(ner[word_idx] if label_all_tokens else -100)94 pos_spacy_ids.append(pos[word_idx] if label_all_tokens else -100)95 dep_spacy_ids.append(dep[word_idx] if label_all_tokens else -100)96 depth_spacy_ids.append(depth[word_idx] if label_all_tokens else -100)97 previous_word_idx = word_idx98 99 ner_spacy.append(ner_spacy_ids)100 pos_spacy.append(pos_spacy_ids)101 dep_spacy.append(dep_spacy_ids)102 depth_spacy.append(depth_spacy_ids)103 104 tokenized_inputs["pos_spacy"] = pos_spacy105 tokenized_inputs["ner_spacy"] = ner_spacy106 tokenized_inputs["dep_spacy"] = dep_spacy107 tokenized_inputs["depth_spacy"] = depth_spacy108 109 return tokenized_inputs110 111 112def find_nearest_nugget_features(doc, start_idx, end_idx, event_nuggets):113 nearest_subtype = None114 nearest_dist = math.inf115 relative_pos = None116 117 mid_idx = (end_idx + start_idx) / 2118 for nugget in event_nuggets:119 mid_nugget_idx = (nugget["nugget"]["startOffset"] + nugget["nugget"]["endOffset"]) / 2120 dist = abs(mid_nugget_idx - mid_idx)121 122 if dist < nearest_dist:123 nearest_dist = dist124 nearest_subtype = nugget["subtype"]125 for sent in doc.sents:126 if between_idxs(mid_idx, sent.start_char, sent.end_char) and between_idxs(mid_nugget_idx, sent.start_char, sent.end_char):127 if mid_idx < mid_nugget_idx:128 relative_pos = "before-same-sentence"129 else:130 relative_pos = "after-same-sentence"131 break132 elif between_idxs(mid_nugget_idx, sent.start_char, sent.end_char) and mid_idx > mid_nugget_idx:133 relative_pos = "after-differ-sentence"134 break135 elif between_idxs(mid_idx, sent.start_char, sent.end_char) and mid_idx < mid_nugget_idx:136 relative_pos = "before-differ-sentence"137 break138 139 nearest_dist = int(min(10, nearest_dist // 20))140 return nearest_subtype, nearest_dist, relative_pos141 142def find_dep_depth(token):143 depth = 0144 current_token = token145 while current_token.head != current_token:146 depth += 1147 current_token = current_token.head148 return min(depth, 16)149 150def between_idxs(idx, start_idx, end_idx):151 return idx >= start_idx and idx <= end_idx