PreranTej/bias-detection-api
0
1import torch2import torch.nn as nn3from transformers import RobertaModel, RobertaPreTrainedModel4 5 6class RobertaMultiTask(RobertaPreTrainedModel):7 def __init__(self, config):8 super().__init__(config)9 self.num_labels = config.num_labels10 self.roberta = RobertaModel(config)11 self.dropout = nn.Dropout(config.hidden_dropout_prob)12 self.classifier = nn.Linear(config.hidden_size, config.num_labels)13 self.span_classifier = nn.Linear(config.hidden_size, 2)14 self.post_init()15 16 def forward(17 self,18 input_ids=None,19 attention_mask=None,20 token_type_ids=None,21 labels=None,22 span_labels=None23 ):24 outputs = self.roberta(25 input_ids,26 attention_mask=attention_mask27 )28 sequence_output = self.dropout(outputs.last_hidden_state)29 pooled_output = self.dropout(outputs.pooler_output)30 31 logits = self.classifier(pooled_output)32 span_logits = self.span_classifier(sequence_output)33 34 loss = None35 if labels is not None and span_labels is not None:36 cls_loss = nn.CrossEntropyLoss()(37 logits.view(-1, self.num_labels),38 labels.view(-1)39 )40 span_loss = nn.CrossEntropyLoss(ignore_index=-100)(41 span_logits.view(-1, 2),42 span_labels.view(-1)43 )44 loss = cls_loss + 0.3 * span_loss45 46 return {47 "loss": loss,48 "logits": logits,49 "span_logits": span_logits50 }