Team Ai
Modelpublic

IEETA/BioNExt-Extractor

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes117downloads
modeling_bionextextractor.py229 linesDownload Raw Back to root
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