Views
No views yet
1import torch
2import torch.nn as nn
3import torch.nn.functional as F
4from torch.nn import CrossEntropyLoss, KLDivLoss
5from transformers.modeling_outputs import TokenClassifierOutput
6from transformers import BertModel, BertPreTrainedModel
7
8class BertForHighlightPrediction(BertPreTrainedModel):
9 _keys_to_ignore_on_load_unexpected = [r"pooler"]
10
11 def __init__(self, config, **model_kwargs):
12 super().__init__(config)
13 # self.model_args = model_kargs["model_args"]
14 self.num_labels = config.num_labels
15 self.bert = BertModel(config, add_pooling_layer=False)
16 classifier_dropout = (
17 config.classifier_dropout if config.classifier_dropout is not None else config.hidden_dropout_prob
18 )
19 self.dropout = nn.Dropout(classifier_dropout)
20 self.tokens_clf = nn.Linear(config.hidden_size, config.num_labels)
21
22 self.tau = model_kwargs.pop('tau', 1)
23 self.gamma = model_kwargs.pop('gamma', 1)
24 self.soft_labeling = model_kwargs.pop('soft_labeling', False)
25
26 self.init_weights()
27 self.softmax = nn.Softmax(dim=-1)
28
29 def forward(self,
30 input_ids=None,
31 probs=None, # soft-labeling
32 attention_mask=None,
33 token_type_ids=None,
34 position_ids=None,
35 head_mask=None,
36 inputs_embeds=None,
37 labels=None,
38 output_attentions=None,
39 output_hidden_states=None,
40 return_dict=None,):
41
42 outputs = self.bert(
43 input_ids,
44 attention_mask=attention_mask,
45 token_type_ids=token_type_ids,
46 position_ids=position_ids,
47 head_mask=head_mask,
48 inputs_embeds=inputs_embeds,
49 output_attentions=output_attentions,
50 output_hidden_states=output_hidden_states,
51 return_dict=return_dict,
52 )
53
54 tokens_output = outputs[0]
55 highlight_logits = self.tokens_clf(self.dropout(tokens_output))
56
57 loss = None
58 if labels is not None:
59 loss_fct = CrossEntropyLoss()
60 active_loss = attention_mask.view(-1) == 1
61 active_logits = highlight_logits.view(-1, self.num_labels)
62 active_labels = torch.where(
63 active_loss,
64 labels.view(-1),
65 torch.tensor(loss_fct.ignore_index).type_as(labels)
66 )
67 loss_ce = loss_fct(active_logits, active_labels)
68
69 loss_kl = 0
70 if self.soft_labeling:
71 loss_fct = KLDivLoss(reduction='sum')
72 active_mask = (attention_mask * token_type_ids).view(-1, 1) # BL 1
73 n_active = (active_mask == 1).sum()
74 active_mask = active_mask.repeat(1, 2) # BL 2
75 input_logp = F.log_softmax(active_logits / self.tau, -1) # BL 2
76 target_p = torch.cat(( (1-probs).view(-1, 1), probs.view(-1, 1)), -1) # BL 2
77
78 loss_kl = loss_fct(input_logp, target_p * active_mask) / n_active
79
80 loss = self.gamma * loss_ce + (1-self.gamma) * loss_kl
81
82 # print("Loss:\n")
83 # print(loss)
84 # print(loss_kl)
85 # print(loss_ce)
86
87 return TokenClassifierOutput(
88 loss=loss,
89 logits=highlight_logits,
90 hidden_states=outputs.hidden_states,
91 attentions=outputs.attentions,
92 )
93 @torch.no_grad()
94 def inference(self, outputs):
95
96 with torch.no_grad():
97 outputs = self.forward(**batch_inputs)
98 probabilities = self.softmax(self.tokens_clf(outputs.hidden_states[-1]))
99 predictions = torch.argmax(probabilities, dim=-1)
100
101 # active filtering
102 active_tokens = batch_inputs['attention_mask'] == 1
103 active_predictions = torch.where(
104 active_tokens,
105 predictions,
106 torch.tensor(-1).type_as(predictions)
107 )
108
109 outputs = {
110 "probabilities": probabilities[:, :, 1].detach(), # shape: (batch, length)
111 "active_predictions": predictions.detach(),
112 "active_tokens": active_tokens,
113 }
114
115 return outputs