(self, sample, batch_idx, _)
| 73 | return self.model |
| 74 | |
| 75 | def _training_step(self, sample, batch_idx, _): |
| 76 | loss_output = self.run_model(self.model, sample) |
| 77 | total_loss = sum([v for v in loss_output.values() if isinstance(v, torch.Tensor) and v.requires_grad]) |
| 78 | loss_output['batch_size'] = sample['txt_tokens'].size()[0] |
| 79 | return total_loss, loss_output |
| 80 | |
| 81 | def validation_step(self, sample, batch_idx): |
| 82 | outputs = {} |