| 157 | return calculate_rouge(preds, target) |
| 158 | |
| 159 | def _generative_step(self, batch: dict) -> dict: |
| 160 | pad_token_id = self.tokenizer.pad_token_id |
| 161 | source_ids, source_mask, y = SummarizationDataset.trim_seq2seq_batch(batch, pad_token_id) |
| 162 | t0 = time.time() |
| 163 | generated_ids = self.model.generate(input_ids=source_ids, attention_mask=source_mask, use_cache=True,) |
| 164 | gen_time = (time.time() - t0) / source_ids.shape[0] |
| 165 | preds = self.ids_to_clean_text(generated_ids) |
| 166 | target = self.ids_to_clean_text(y) |
| 167 | loss_tensors = self._step(batch) |
| 168 | base_metrics = {name: loss for name, loss in zip(self.loss_names, loss_tensors)} |
| 169 | rouge: Dict = self.calc_generative_metrics(preds, target) |
| 170 | summ_len = np.mean(lmap(len, generated_ids)) |
| 171 | base_metrics.update(gen_time=gen_time, summ_len=summ_len, preds=preds, target=target, **rouge) |
| 172 | return base_metrics |
| 173 | |
| 174 | def test_step(self, batch, batch_idx): |
| 175 | return self._generative_step(batch) |