r""" :class:`transformers.LogitsProcessor` that enforces no repetition of n-grams. See `Fairseq `__. Args: ngram_size (:obj:`int`): All ngrams of size
| 442 | |
| 443 | |
| 444 | class NoRepeatNGramLogitsProcessor(LogitsProcessor): |
| 445 | r""" |
| 446 | :class:`transformers.LogitsProcessor` that enforces no repetition of n-grams. See `Fairseq |
| 447 | <https://github.com/pytorch/fairseq/blob/a07cb6f40480928c9e0548b737aadd36ee66ac76/fairseq/sequence_generator.py#L345>`__. |
| 448 | |
| 449 | Args: |
| 450 | ngram_size (:obj:`int`): |
| 451 | All ngrams of size :obj:`ngram_size` can only occur once. |
| 452 | """ |
| 453 | |
| 454 | def __init__(self, ngram_size: int): |
| 455 | if not isinstance(ngram_size, int) or ngram_size <= 0: |
| 456 | raise ValueError(f"`ngram_size` has to be a strictly positive integer, but is {ngram_size}") |
| 457 | self.ngram_size = ngram_size |
| 458 | |
| 459 | def __call__(self, input_ids: torch.LongTensor, scores: torch.FloatTensor) -> torch.FloatTensor: |
| 460 | num_batch_hypotheses = scores.shape[0] |
| 461 | cur_len = input_ids.shape[-1] |
| 462 | banned_batch_tokens = self._calc_banned_ngram_tokens(input_ids, num_batch_hypotheses, cur_len) |
| 463 | |
| 464 | for i, banned_tokens in enumerate(banned_batch_tokens): |
| 465 | scores[i, banned_tokens] = -float("inf") |
| 466 | |
| 467 | return scores |
| 468 | |
| 469 | def _calc_banned_ngram_tokens( |
| 470 | self, prev_input_ids: torch.Tensor, num_hypos: int, cur_len: int |
| 471 | ) -> List[Iterable[int]]: |
| 472 | """Copied from fairseq for no_repeat_ngram in beam_search""" |
| 473 | if cur_len + 1 < self.ngram_size: |
| 474 | # return no banned tokens if we haven't generated no_repeat_ngram_size tokens yet |
| 475 | return [[] for _ in range(num_hypos)] |
| 476 | generated_ngrams = [{} for _ in range(num_hypos)] |
| 477 | for idx in range(num_hypos): |
| 478 | gen_tokens = prev_input_ids[idx].tolist() |
| 479 | generated_ngram = generated_ngrams[idx] |
| 480 | for ngram in zip(*[gen_tokens[i:] for i in range(self.ngram_size)]): |
| 481 | prev_ngram_tuple = tuple(ngram[:-1]) |
| 482 | generated_ngram[prev_ngram_tuple] = generated_ngram.get(prev_ngram_tuple, []) + [ngram[-1]] |
| 483 | |
| 484 | def _get_generated_ngrams(hypo_idx): |
| 485 | # Before decoding the next token, prevent decoding of ngrams that have already appeared |
| 486 | start_idx = cur_len + 1 - self.ngram_size |
| 487 | ngram_idx = tuple(prev_input_ids[hypo_idx, start_idx:cur_len].tolist()) |
| 488 | return generated_ngrams[hypo_idx].get(ngram_idx, []) |
| 489 | |
| 490 | banned_tokens = [_get_generated_ngrams(hypo_idx) for hypo_idx in range(num_hypos)] |
| 491 | return banned_tokens |