Team Ai
Modelpublic

CyberPeace-Institute/Cybersecurity-Knowledge-Graph

sourceHugging Facemitupdated 3y agoView on Hugging Face
23likes60downloads
nugget_model_utils.py151 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"]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