(self, *input, **kwargs)
| 34 | |
| 35 | |
| 36 | def forward(self, *input, **kwargs): |
| 37 | input_ids = kwargs.pop("input_ids") |
| 38 | |
| 39 | pad_token_id = kwargs.pop("pad_token_id") |
| 40 | attention_mask = (input_ids != pad_token_id).long() |
| 41 | |
| 42 | if self.training: |
| 43 | task = kwargs.pop("task") |
| 44 | if task == "mlm": |
| 45 | output_ids = kwargs.pop('labels') |
| 46 | y_ids = output_ids[:, :-1].contiguous() |
| 47 | lm_labels = output_ids[:, 1:].clone() |
| 48 | lm_labels[output_ids[:, 1:] == pad_token_id] = -100 |
| 49 | |
| 50 | outputs = self.bert(input_ids, |
| 51 | attention_mask=attention_mask, decoder_input_ids=y_ids, lm_labels=lm_labels, ) |
| 52 | return (outputs[0],) |
| 53 | elif task == "col_pred": |
| 54 | label_ids = kwargs.pop("labels") |
| 55 | column_spans = kwargs.pop("column_spans") |
| 56 | column_selection_prob = self.column_prediction(input_ids, attention_mask, column_spans) |
| 57 | label_mask = column_spans.view(-1, 2)[:,0] > 0 |
| 58 | |
| 59 | column_selection_loss = F.binary_cross_entropy(column_selection_prob.view(-1)[label_mask], label_ids.view(-1)[label_mask].float(), |
| 60 | reduction="sum") / label_ids.size(0) |
| 61 | return (column_selection_loss, ) |
| 62 | else: |
| 63 | raise NotImplementedError("Unknown task {}".format(task)) |
| 64 | |
| 65 | else: |
| 66 | task = kwargs.pop("task") |
| 67 | |
| 68 | if task == "mlm": |
| 69 | label_eos_id = kwargs.pop("label_eos_id") |
| 70 | label_bos_id = kwargs.pop("label_bos_id") |
| 71 | label_padding_id = kwargs.pop("label_padding_id") |
| 72 | generated_ids = self.bert.generate( |
| 73 | input_ids=input_ids, |
| 74 | attention_mask=attention_mask, |
| 75 | num_beams=3, |
| 76 | max_length=input_ids.size(1) + 5, |
| 77 | length_penalty=2.0, |
| 78 | early_stopping=True, |
| 79 | use_cache=True, |
| 80 | decoder_start_token_id=label_bos_id, |
| 81 | eos_token_id=label_eos_id, |
| 82 | pad_token_id=label_padding_id |
| 83 | ) |
| 84 | |
| 85 | output_ids = kwargs.pop('labels') |
| 86 | y_ids = output_ids[:, :-1].contiguous() |
| 87 | lm_labels = output_ids[:, 1:].clone() |
| 88 | lm_labels[output_ids[:, 1:] == pad_token_id] = -100 |
| 89 | |
| 90 | outputs = self.bert(input_ids, |
| 91 | attention_mask=attention_mask, decoder_input_ids=y_ids, lm_labels=lm_labels, ) |
| 92 | |
| 93 | return (outputs[0].detach(), generated_ids) |
nothing calls this directly
no test coverage detected