| 203 | |
| 204 | |
| 205 | class DecoderClassifier(nn.Module): |
| 206 | def __init__(self, encoder, config, tokenizer, args): |
| 207 | super(DecoderClassifier, self).__init__() |
| 208 | self.encoder = encoder |
| 209 | self.config=config |
| 210 | self.tokenizer=tokenizer |
| 211 | self.args=args |
| 212 | self.classifier = nn.Linear(config.hidden_size, 2) |
| 213 | |
| 214 | def forward(self, input_ids=None, labels=None, weight=None): |
| 215 | attention_mask = input_ids.ne(self.tokenizer.pad_token_id) |
| 216 | outputs = self.encoder(input_ids, attention_mask=attention_mask) |
| 217 | hidden_states = outputs[0] |
| 218 | logits = self.classifier(hidden_states) |
| 219 | |
| 220 | batch_size = input_ids.size(0) |
| 221 | if self.config.pad_token_id is None and batch_size != 1: |
| 222 | raise ValueError("Cannot handle batch sizes > 1 if no padding token is defined.") |
| 223 | if self.config.pad_token_id is None: |
| 224 | sequence_lengths = -1 |
| 225 | else: |
| 226 | if input_ids is not None: |
| 227 | # if no pad token found, use modulo instead of reverse indexing for ONNX compatibility |
| 228 | sequence_lengths = torch.eq(input_ids, self.config.pad_token_id).int().argmax(-1) - 1 |
| 229 | sequence_lengths = sequence_lengths % input_ids.shape[-1] |
| 230 | sequence_lengths = sequence_lengths.to(logits.device) |
| 231 | else: |
| 232 | sequence_lengths = -1 |
| 233 | |
| 234 | pooled_logits = logits[torch.arange(batch_size, device=logits.device), sequence_lengths] |
| 235 | prob = nn.functional.softmax(pooled_logits, dim=-1) |
| 236 | |
| 237 | |
| 238 | if labels is not None: |
| 239 | labels = labels.to(logits.device) |
| 240 | loss_fct = nn.CrossEntropyLoss(weight=weight) |
| 241 | |
| 242 | loss = loss_fct(pooled_logits.view(-1, 2), labels.view(-1)) |
| 243 | return loss, prob |
| 244 | else: |
| 245 | return prob |