MCPcopy Create free account
hub / github.com/THUDM/GLM / NoRepeatNGramLogitsProcessor

Class NoRepeatNGramLogitsProcessor

generation_utils.py:444–491  ·  view source on GitHub ↗

r""" :class:`transformers.LogitsProcessor` that enforces no repetition of n-grams. See `Fairseq `__. Args: ngram_size (:obj:`int`): All ngrams of size

Source from the content-addressed store, hash-verified

442
443
444class 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

Callers 1

__init__Method · 0.90

Calls

no outgoing calls

Tested by

no test coverage detected