PreranTej commited on
Commit
358e141
·
verified ·
1 Parent(s): 1d31d28

Update modeling_roberta_multitask.py

Browse files
Files changed (1) hide show
  1. modeling_roberta_multitask.py +30 -10
modeling_roberta_multitask.py CHANGED
@@ -1,8 +1,8 @@
1
- # Copy exact same file from your HuggingFace model repo
2
  import torch
3
  import torch.nn as nn
4
  from transformers import RobertaModel, RobertaPreTrainedModel
5
 
 
6
  class RobertaMultiTask(RobertaPreTrainedModel):
7
  def __init__(self, config):
8
  super().__init__(config)
@@ -13,18 +13,38 @@ class RobertaMultiTask(RobertaPreTrainedModel):
13
  self.span_classifier = nn.Linear(config.hidden_size, 2)
14
  self.post_init()
15
 
16
- def forward(self, input_ids=None, attention_mask=None,
17
- token_type_ids=None, labels=None, span_labels=None):
18
- outputs = self.roberta(input_ids, attention_mask=attention_mask)
 
 
 
 
 
 
 
 
 
19
  sequence_output = self.dropout(outputs.last_hidden_state)
20
  pooled_output = self.dropout(outputs.pooler_output)
21
- logits = self.classifier(pooled_output)
22
- span_logits = self.span_classifier(sequence_output)
 
 
23
  loss = None
24
  if labels is not None and span_labels is not None:
25
- cls_loss = nn.CrossEntropyLoss()(
26
- logits.view(-1, self.num_labels), labels.view(-1))
 
 
27
  span_loss = nn.CrossEntropyLoss(ignore_index=-100)(
28
- span_logits.view(-1, 2), span_labels.view(-1))
 
 
29
  loss = cls_loss + 0.3 * span_loss
30
- return {"loss": loss, "logits": logits, "span_logits": span_logits}
 
 
 
 
 
 
 
1
  import torch
2
  import torch.nn as nn
3
  from transformers import RobertaModel, RobertaPreTrainedModel
4
 
5
+
6
  class RobertaMultiTask(RobertaPreTrainedModel):
7
  def __init__(self, config):
8
  super().__init__(config)
 
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=None
23
+ ):
24
+ outputs = self.roberta(
25
+ input_ids,
26
+ attention_mask=attention_mask
27
+ )
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 = None
35
  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_loss
45
+
46
+ return {
47
+ "loss": loss,
48
+ "logits": logits,
49
+ "span_logits": span_logits
50
+ }