MCPcopy Create free account
hub / github.com/huggingface/transformers / _generative_step

Method _generative_step

examples/seq2seq/finetune.py:159–172  ·  view source on GitHub ↗
(self, batch: dict)

Source from the content-addressed store, hash-verified

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)

Callers 2

validation_stepMethod · 0.95
test_stepMethod · 0.95

Calls 7

ids_to_clean_textMethod · 0.95
_stepMethod · 0.95
lmapFunction · 0.90
trim_seq2seq_batchMethod · 0.80
updateMethod · 0.80
generateMethod · 0.45

Tested by 1

test_stepMethod · 0.76