Team Ai
Modelpublic

HCKLab/BiBert-MultiTask-1

sourceHugging Facemitupdated 4y agoView on Hugging Face
0likes6downloads
bert_for_sequence_classification.py146 linesDownload Raw Back to root
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