Team Ai
Modelpublic

CyberPeace-Institute/Cybersecurity-Knowledge-Graph

sourceHugging Facemitupdated 3y agoView on Hugging Face
23likes60downloads
args_model_utils.py210 linesDownload Raw Back to root
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"]21dep_spacy_tag_list = list(nlp.get_pipe("parser").labels)22event_nugget_tag_list = ["Databreach", "Ransom", "PatchVulnerability", "Phishing", "DiscoverVulnerability"]23arg_nugget_relative_pos_tag_list = ["before-same-sentence", "before-differ-sentence", "after-same-sentence", "after-differ-sentence"]24 25class CustomRobertaWithPOS(nn.Module):26    def __init__(self, num_classes):27        super(CustomRobertaWithPOS, self).__init__()28        self.num_classes = num_classes29 30        self.pos_embed = nn.Embedding(len(pos_spacy_tag_list), 16)31        self.ner_embed = nn.Embedding(len(ner_spacy_tag_list), 8)32        self.dep_embed = nn.Embedding(len(dep_spacy_tag_list), 8)33        self.depth_embed = nn.Embedding(17, 8)34        self.subtype_embed = nn.Embedding(len(event_nugget_tag_list), 2)35        self.dist_embed = nn.Embedding(11, 6)36        self.relative_pos_embed = nn.Embedding(len(arg_nugget_relative_pos_tag_list), 2)37 38        self.roberta = roberta_model39        self.dropout1 = nn.Dropout(0.2)40        self.fc1 = nn.Linear(self.roberta.config.hidden_size + 50, num_classes)41 42    def forward(self, input_ids, attention_mask, pos_spacy, ner_spacy, dep_spacy, depth_spacy, nearest_nugget_subtype, nearest_nugget_dist, arg_nugget_relative_pos):43        outputs = self.roberta(input_ids=input_ids, attention_mask=attention_mask)44        last_hidden_output = outputs.last_hidden_state45        46        pooler_output = outputs.pooler_output47        pooler_output_unsqz = pooler_output.unsqueeze(1)48        pooler_output_fin = pooler_output_unsqz.expand(-1, last_hidden_output.shape[1], -1)49 50 51        pos_mask = pos_spacy != -10052        pos_embed_masked = self.pos_embed(pos_spacy[pos_mask])53        pos_embed = torch.zeros((pos_spacy.shape[0], pos_spacy.shape[1], 16), dtype=torch.float).to(device)54        pos_embed[pos_mask] = pos_embed_masked55 56        ner_mask = ner_spacy != -10057        ner_embed_masked = self.ner_embed(ner_spacy[ner_mask])58        ner_embed = torch.zeros((ner_spacy.shape[0], ner_spacy.shape[1], 8), dtype=torch.float).to(device)59        ner_embed[ner_mask] = ner_embed_masked60 61        dep_mask = dep_spacy != -10062        dep_embed_masked = self.dep_embed(dep_spacy[dep_mask])63        dep_embed = torch.zeros((dep_spacy.shape[0], dep_spacy.shape[1], 8), dtype=torch.float).to(device)64        dep_embed[dep_mask] = dep_embed_masked65 66        depth_mask = depth_spacy != -10067        depth_embed_masked = self.depth_embed(depth_spacy[depth_mask])68        depth_embed = torch.zeros((depth_spacy.shape[0], depth_spacy.shape[1], 8), dtype=torch.float).to(device)69        depth_embed[dep_mask] = depth_embed_masked70 71        nearest_nugget_subtype_mask = nearest_nugget_subtype != -10072        nearest_nugget_subtype_embed_masked = self.subtype_embed(nearest_nugget_subtype[nearest_nugget_subtype_mask])73        nearest_nugget_subtype_embed = torch.zeros((nearest_nugget_subtype.shape[0], nearest_nugget_subtype.shape[1], 2), dtype=torch.float).to(device)74        nearest_nugget_subtype_embed[dep_mask] = nearest_nugget_subtype_embed_masked75 76        nearest_nugget_dist_mask = nearest_nugget_dist != -10077        nearest_nugget_dist_embed_masked = self.dist_embed(nearest_nugget_dist[nearest_nugget_dist_mask])78        nearest_nugget_dist_embed = torch.zeros((nearest_nugget_dist.shape[0], nearest_nugget_dist.shape[1], 6), dtype=torch.float).to(device)79        nearest_nugget_dist_embed[dep_mask] = nearest_nugget_dist_embed_masked80 81        arg_nugget_relative_pos_mask = arg_nugget_relative_pos != -10082        arg_nugget_relative_pos_embed_masked = self.relative_pos_embed(arg_nugget_relative_pos[arg_nugget_relative_pos_mask])83        arg_nugget_relative_pos_embed = torch.zeros((arg_nugget_relative_pos.shape[0], arg_nugget_relative_pos.shape[1], 2), dtype=torch.float).to(device)84        arg_nugget_relative_pos_embed[dep_mask] = arg_nugget_relative_pos_embed_masked85 86        features_concat = torch.cat((last_hidden_output, pos_embed, ner_embed, dep_embed, depth_embed, nearest_nugget_subtype_embed, nearest_nugget_dist_embed, arg_nugget_relative_pos_embed), 2).to(device)87        features_concat = self.dropout1(features_concat)88 89        logits = self.fc1(features_concat)90 91        return logits92 93 94def tokenize_and_align_labels_with_pos_ner_dep(examples, tokenizer, label_all_tokens = True):95    tokenized_inputs = tokenizer(examples["tokens"], padding='max_length', truncation=True, is_split_into_words=True)96    #tokenized_inputs.pop('input_ids')97    ner_spacy = []98    pos_spacy = []99    dep_spacy = []100    depth_spacy = []101    nearest_nugget_subtype = []102    nearest_nugget_dist = []103    arg_nugget_relative_pos = []104 105    for i, (pos, ner, dep, depth, subtype, dist, relative_pos) in enumerate(zip(examples["pos_spacy"], 106                                                                                examples["ner_spacy"], 107                                                                                examples["dep_spacy"], 108                                                                                examples["depth_spacy"], 109                                                                                examples["nearest_nugget_subtype"], 110                                                                                examples["nearest_nugget_dist"], 111                                                                                examples["arg_nugget_relative_pos"])):112        word_ids = tokenized_inputs.word_ids(batch_index=i)113        previous_word_idx = None114        ner_spacy_ids = []115        pos_spacy_ids = []116        dep_spacy_ids = []117        depth_spacy_ids = []118        nearest_nugget_subtype_ids = []119        nearest_nugget_dist_ids = []120        arg_nugget_relative_pos_ids = []121 122        for word_idx in word_ids:123            # Special tokens have a word id that is None. We set the label to -100 so they are automatically124            # ignored in the loss function.125            if word_idx is None:126                ner_spacy_ids.append(-100)127                pos_spacy_ids.append(-100)128                dep_spacy_ids.append(-100)129                depth_spacy_ids.append(-100)130                nearest_nugget_subtype_ids.append(-100)131                nearest_nugget_dist_ids.append(-100)132                arg_nugget_relative_pos_ids.append(-100)133            # We set the label for the first token of each word.134            elif word_idx != previous_word_idx:135                ner_spacy_ids.append(ner[word_idx])136                pos_spacy_ids.append(pos[word_idx])137                dep_spacy_ids.append(dep[word_idx])138                depth_spacy_ids.append(depth[word_idx])139                nearest_nugget_subtype_ids.append(subtype[word_idx])140                nearest_nugget_dist_ids.append(dist[word_idx])141                arg_nugget_relative_pos_ids.append(relative_pos[word_idx])142            # For the other tokens in a word, we set the label to either the current label or -100, depending on143            # the label_all_tokens flag.144            else:145                ner_spacy_ids.append(ner[word_idx] if label_all_tokens else -100)146                pos_spacy_ids.append(pos[word_idx] if label_all_tokens else -100)147                dep_spacy_ids.append(dep[word_idx] if label_all_tokens else -100)148                depth_spacy_ids.append(depth[word_idx] if label_all_tokens else -100)149                nearest_nugget_subtype_ids.append(subtype[word_idx] if label_all_tokens else -100)150                nearest_nugget_dist_ids.append(dist[word_idx] if label_all_tokens else -100)151                arg_nugget_relative_pos_ids.append(relative_pos[word_idx] if label_all_tokens else -100)152            previous_word_idx = word_idx153 154        ner_spacy.append(ner_spacy_ids)155        pos_spacy.append(pos_spacy_ids)156        dep_spacy.append(dep_spacy_ids)157        depth_spacy.append(depth_spacy_ids)158        nearest_nugget_subtype.append(nearest_nugget_subtype_ids)159        nearest_nugget_dist.append(nearest_nugget_dist_ids)160        arg_nugget_relative_pos.append(arg_nugget_relative_pos_ids)161 162    tokenized_inputs["pos_spacy"] = pos_spacy163    tokenized_inputs["ner_spacy"] = ner_spacy164    tokenized_inputs["dep_spacy"] = dep_spacy165    tokenized_inputs["depth_spacy"] = depth_spacy166    tokenized_inputs["nearest_nugget_subtype"] = nearest_nugget_subtype167    tokenized_inputs["nearest_nugget_dist"] = nearest_nugget_dist168    tokenized_inputs["arg_nugget_relative_pos"] = arg_nugget_relative_pos 169    return tokenized_inputs170 171def find_nearest_nugget_features(doc, start_idx, end_idx, event_nuggets):172            nearest_subtype = None173            nearest_dist = math.inf174            relative_pos = None175 176            mid_idx = (end_idx + start_idx) / 2177            for nugget in event_nuggets:178                mid_nugget_idx = (nugget["startOffset"] + nugget["endOffset"]) / 2179                dist = abs(mid_nugget_idx - mid_idx)180 181                if dist < nearest_dist:182                    nearest_dist = dist183                    nearest_subtype = nugget["subtype"]184                    for sent in doc.sents:185                        if between_idxs(mid_idx, sent.start_char, sent.end_char) and between_idxs(mid_nugget_idx, sent.start_char, sent.end_char):186                            if mid_idx < mid_nugget_idx:187                                relative_pos = "before-same-sentence"188                            else:189                                relative_pos = "after-same-sentence"190                            break191                        elif between_idxs(mid_nugget_idx, sent.start_char, sent.end_char) and mid_idx > mid_nugget_idx:192                            relative_pos = "after-differ-sentence"193                            break194                        elif between_idxs(mid_idx, sent.start_char, sent.end_char) and mid_idx < mid_nugget_idx:195                            relative_pos = "before-differ-sentence"196                            break197            198            nearest_dist = int(min(10, nearest_dist // 20))199            return nearest_subtype, nearest_dist, relative_pos200 201def find_dep_depth(token):202            depth = 0203            current_token = token204            while current_token.head != current_token:205                depth += 1206                current_token = current_token.head207            return min(depth, 16)208        209def between_idxs(idx, start_idx, end_idx):210    return idx >= start_idx and idx <= end_idx