HCKLab/BiBert-MultiTask-1
06
1import torch2import transformers3from torch import nn4from torch.nn import BCEWithLogitsLoss, CrossEntropyLoss, MSELoss5from typing import List, Optional, Tuple, Union6 7from transformers import BertTokenizer8from transformers import models, DataCollatorWithPadding, AutoTokenizer9from transformers.modeling_outputs import SequenceClassifierOutput10 11from transformers.models.bert.configuration_bert import BertConfig12from transformers.models.bert.modeling_bert import (13 BertPreTrainedModel,14 BERT_INPUTS_DOCSTRING,15 _TOKENIZER_FOR_DOC,16 _CHECKPOINT_FOR_DOC,17 BERT_START_DOCSTRING,18 _CONFIG_FOR_DOC,19 _SEQ_CLASS_EXPECTED_OUTPUT,20 _SEQ_CLASS_EXPECTED_LOSS,21 BertModel,22)23 24from transformers.file_utils import (25 add_code_sample_docstrings,26 add_start_docstrings_to_model_forward,27 add_start_docstrings28)29 30@add_start_docstrings(31 """32 Bert Model transformer with a sequence classification/regression head on top (a linear layer on top of the pooled33 output) e.g. for GLUE tasks.34 """,35 BERT_START_DOCSTRING,36)37class BertForSequenceClassification(BertPreTrainedModel):38 def __init__(self, config, **kwargs):39 super().__init__(transformers.PretrainedConfig())40 #task_labels_map={"binary_classification": 2, "label_classification": 5}41 self.tasks = kwargs.get("tasks_map", {})42 self.config = config43 44 self.bert = BertModel(config)45 classifier_dropout = (46 config.classifier_dropout47 if config.classifier_dropout is not None48 else config.hidden_dropout_prob49 )50 self.dropout = nn.Dropout(classifier_dropout)51 ## add task specific output heads52 self.classifier1 = nn.Linear(53 config.hidden_size, self.tasks[0].num_labels54 )55 self.classifier2 = nn.Linear(56 config.hidden_size, self.tasks[1].num_labels57 )58 59 self.init_weights()60 61 @add_start_docstrings_to_model_forward(62 BERT_INPUTS_DOCSTRING.format("batch_size, sequence_length")63 )64 @add_code_sample_docstrings(65 processor_class=_TOKENIZER_FOR_DOC,66 checkpoint=_CHECKPOINT_FOR_DOC,67 output_type=SequenceClassifierOutput,68 config_class=_CONFIG_FOR_DOC,69 expected_output=_SEQ_CLASS_EXPECTED_OUTPUT,70 expected_loss=_SEQ_CLASS_EXPECTED_LOSS,71 )72 def forward(73 self,74 input_ids: Optional[torch.Tensor] = None,75 attention_mask: Optional[torch.Tensor] = None,76 token_type_ids: Optional[torch.Tensor] = None,77 position_ids: Optional[torch.Tensor] = None,78 head_mask: Optional[torch.Tensor] = None,79 inputs_embeds: Optional[torch.Tensor] = None,80 labels: Optional[torch.Tensor] = None,81 output_attentions: Optional[bool] = None,82 output_hidden_states: Optional[bool] = None,83 return_dict: Optional[bool] = None,84 task_ids=None,85) -> Union[Tuple[torch.Tensor], SequenceClassifierOutput]:86 r"""87 labels (:obj:`torch.LongTensor` of shape :obj:`(batch_size,)`, `optional`):88 Labels for computing the sequence classification/regression loss. Indices should be in :obj:`[0, ...,89 config.num_labels - 1]`. If :obj:`config.num_labels == 1` a regression loss is computed (Mean-Square loss),90 If :obj:`config.num_labels > 1` a classification loss is computed (Cross-Entropy).91 """92 return_dict = (93 return_dict if return_dict is not None else self.config.use_return_dict94 )95 96 outputs = self.bert(97 input_ids,98 attention_mask=attention_mask,99 token_type_ids=token_type_ids,100 position_ids=position_ids,101 head_mask=head_mask,102 inputs_embeds=inputs_embeds,103 output_attentions=output_attentions,104 output_hidden_states=output_hidden_states,105 return_dict=return_dict,106 )107 108 pooled_output = outputs[1]109 110 pooled_output = self.dropout(pooled_output)111 112 unique_task_ids_list = torch.unique(task_ids).tolist()113 loss_list = []114 logits = None115 for unique_task_id in unique_task_ids_list:116 117 loss = None118 task_id_filter = task_ids == unique_task_id 119 120 if unique_task_id == 0:121 logits = self.classifier1(pooled_output[task_id_filter])122 elif unique_task_id == 1:123 logits = self.classifier2(pooled_output[task_id_filter])124 125 126 if labels is not None: 127 loss_fct = CrossEntropyLoss()128 loss = loss_fct(logits.view(-1, self.tasks[unique_task_id].num_labels), labels[task_id_filter].view(-1))129 loss_list.append(loss)130 131 # logits are only used for eval. and in case of eval the batch is not multi task132 # For training only the loss is used133 134 if loss_list:135 loss = torch.stack(loss_list).mean()136 if not return_dict:137 output = (logits,) + outputs[2:]138 return ((loss,) + output) if loss is not None else output139 140 return SequenceClassifierOutput(141 loss=loss,142 logits=logits,143 hidden_states=outputs.hidden_states,144 attentions=outputs.attentions,145 )146 