IEETA/BioNExt-Extractor
0117
1 2import os3from typing import Optional, Union4from transformers import BertModel, PreTrainedModel, AutoConfig, BertModel5from transformers.modeling_outputs import TokenClassifierOutput6from torch import nn7from torch.nn import CrossEntropyLoss8 9from typing import List, Optional10 11import torch12from itertools import islice13from .configuration_bionextextractor import BioNExtExtractorConfig14 15 16import torch17 18from transformers import AutoModel, PreTrainedModel, AutoConfig, BertConfig19from transformers.modeling_outputs import TokenClassifierOutput, SequenceClassifierOutput20 21from torch.nn import CrossEntropyLoss22import math23 24class RelationLossMixin:25 26 def model_loss(self, logits, labels, novel=None, reduction=None):27 if reduction is None:28 return torch.nn.functional.cross_entropy(logits.view(-1, self.num_labels), labels.view(-1))29 else:30 return torch.nn.functional.cross_entropy(logits.view(-1,self.num_labels), labels.view(-1), reduction=reduction)31 32class RelationAndNovelLossMixin(RelationLossMixin):33 34 def model_loss(self, logits, labels, novel=None):35 relation_logits, novel_logits = logits36 relation_loss = super().model_loss(relation_logits, labels, reduction="none")37 novel_loss = torch.nn.functional.cross_entropy(novel_logits.view(-1, 2), novel.view(-1), reduction="none")38 per_sample_loss = relation_loss + (labels!=8).type(logits[0].dtype)*novel_loss39 40 return per_sample_loss.mean()#relation_loss + (labels!=8).type(logits[0].dtype)*novel_loss(novel_logits.view(-1, 2), novel.view(-1))41 #return relation_loss + novel_loss(novel_logits.view(-1, 2), novel.view(-1))42 43class RelationClassifierBase(PreTrainedModel, RelationLossMixin):44 #_keys_to_ignore_on_load_unexpected = [r"pooler"]45 config_class=BioNExtExtractorConfig46 47 def __init__(self, config):48 super().__init__(config)49 self.num_labels = config.num_labels50 self.config = config51 #print(config)52 self.bert = BertModel(config, add_pooling_layer=False)53 54 def training_mode(self):55 if self.config.update_vocab is not None:56 self.bert.resize_token_embeddings(self.config.update_vocab)57 58 def group_embeddings_by_index(self, embeddings, indexes):59 assert len(embeddings.shape)==360 61 batch_size = indexes.shape[0]62 max_tokens = embeddings.shape[1]63 emb_size = embeddings.shape[2]64 # masking padding65 mask_index = indexes!=-166 67 # convert index to 1d of valid index (ignore paddings)68 indexes = indexes + mask_index*(torch.arange(batch_size).to(self.device)*max_tokens).view(batch_size,1,1)69 indexes = indexes.masked_select(mask_index)70 71 # reshape 72 embeddings = embeddings.view(batch_size*max_tokens, emb_size)73 74 # get the embeddings by index75 selected_embeddings_by_index = torch.index_select(embeddings, 0, indexes)76 77 final_output_shape = (mask_index.shape[0], mask_index.shape[1], emb_size)78 group_embeddings = torch.zeros(final_output_shape, dtype=embeddings.dtype).to(self.device).masked_scatter(mask_index, selected_embeddings_by_index)79 80 return group_embeddings, mask_index81 82 def classifier_representation(self, embeddings, mask = None):83 raise NotImplementedError("This is base class, pleas extend an implement classifier_representation")84 85 def classifier(self, class_representation, relation_mask = None):86 raise NotImplementedError("This is base class, pleas extend an implement classifier")87 88 def forward(self,89 input_ids,90 indexes=None,91 novel=None,92 labels=None,93 mask=None,94 return_dict=None,95 **model_kwargs96 ):97 # Default `model.config.use_return_dict´ is `True´98 return_dict = return_dict if return_dict is not None else self.config.use_return_dict99 100 outputs = self.bert(input_ids, return_dict=return_dict, **model_kwargs)101 102 assert indexes is not None103 104 embeddings = outputs.last_hidden_state105 106 selected_embeddings, mask_group = self.group_embeddings_by_index(embeddings, indexes)107 108 class_representation = self.classifier_representation(selected_embeddings, mask_group)109 110 logits = self.classifier(class_representation, relation_mask=mask)111 112 loss = None113 if labels is not None:114 loss = self.model_loss(logits, labels, novel)115 116 117 return SequenceClassifierOutput(118 loss=loss,119 logits=logits,120 hidden_states=outputs.hidden_states,121 attentions=outputs.attentions,122 )123 124 125class RelationClassifierBiLSTM(RelationClassifierBase):126 127 def __init__(self, config):128 super().__init__(config)129 self.num_lstm_layers = config.num_lstm_layers130 self.lstm = torch.nn.LSTM(config.hidden_size, (config.hidden_size) // 2, self.num_lstm_layers, batch_first=True, bidirectional=True)131 self.fc = torch.nn.Linear(config.hidden_size, self.num_labels) # 2 for bidirection132 133 def training_mode(self):134 super().training_mode()135 self.lstm.reset_parameters()136 self.fc.reset_parameters()137 138 def classifier_representation(self, embeddings, mask=None):139 out, _ = self.lstm(embeddings)140 return out[:, -1, :]141 142 def classifier(self, class_representation, mask=None):143 return self.fc(class_representation)144 145class RelationAndNovelClassifierBiLSTM(RelationClassifierBiLSTM, RelationAndNovelLossMixin):146 147 def __init__(self, config):148 super().__init__(config)149 self.fc_novel = torch.nn.Linear(config.hidden_size, 2) # 2 for bidirection150 151 def training_mode(self):152 super().training_mode()153 self.fc_novel.reset_parameters()154 155 def classifier(self, class_representation):156 return super().classifier(class_representation), self.fc_novel(class_representation)157 158class RelationClassifierMHAttention(RelationClassifierBase):159 160 def __init__(self, config):161 super().__init__(config)162 163 self.weight = torch.nn.Parameter(torch.Tensor(1,1,config.hidden_size))164 torch.nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))165 166 self.MHattention_layer = torch.nn.MultiheadAttention(config.hidden_size, config.num_attention_heads, batch_first=True) # 2 for bidirection167 self.fc1 = torch.nn.Linear(config.hidden_size, config.hidden_size//2) # 2 for bidirection168 self.fc1_activation = torch.nn.GELU(approximate='none')169 self.fc2 = torch.nn.Linear(config.hidden_size//2, self.num_labels) # 2 for bidirection170 171 def training_mode(self):172 super().training_mode()173 torch.nn.init.kaiming_uniform_(self.weight, a=math.sqrt(5))174 self.MHattention_layer._reset_parameters()175 self.fc1.reset_parameters()176 self.fc2.reset_parameters()177 178 def classifier_representation(self, embeddings, mask=None):179 batch_size = embeddings.shape[0]180 weight = self.weight.repeat(batch_size, 1, 1)181 182 if mask is not None:183 # flip184 mask = mask.squeeze(-1)==False185 186 out_tensors, _ = self.MHattention_layer(weight, embeddings, embeddings, key_padding_mask=mask)187 188 return out_tensors189 190 def classifier(self, class_representation, relation_mask = None):191 192 x = self.fc1(class_representation)193 x = self.fc1_activation(x) 194 logits = self.fc2(x)195 if relation_mask is not None:196 #print(logits.shape, relation_mask.shape)197 logits = logits + relation_mask.view(-1,1,self.num_labels)198 return logits199 200class RelationAndNovelClassifierMHAttention(RelationClassifierMHAttention, RelationAndNovelLossMixin):201 def __init__(self, config):202 super().__init__(config)203 204 self.fc1_novel = torch.nn.Linear(config.hidden_size, config.hidden_size//2) # 2 for bidirection205 self.fc1_novel_activation = torch.nn.GELU(approximate='none')206 self.fc2_novel = torch.nn.Linear(config.hidden_size//2, 2) # 2 for bidirection207 208 def training_mode(self):209 super().training_mode()210 self.fc1_novel.reset_parameters()211 self.fc2_novel.reset_parameters()212 213 def classifier(self, class_representation, relation_mask=None):214 x = self.fc1_novel(class_representation)215 x = self.fc1_novel_activation(x)216 217 return super().classifier(class_representation, relation_mask=relation_mask), self.fc2_novel(x)218 219ARCH_MAPPING = {"mhawNovelty": RelationAndNovelClassifierMHAttention, 220 "mha": RelationClassifierMHAttention,221 "bilstmwNovelty" : RelationAndNovelClassifierBiLSTM,222 "bilstm": RelationClassifierBiLSTM}223 224## Changing the name to be compatible with HF API225 226class BioNExtExtractorModel(RelationAndNovelClassifierMHAttention):227 config_class=BioNExtExtractorConfig228 229 