MCPcopy Create free account
hub / github.com/SooLab/CGFormer / _get_generated_ngrams

Function _get_generated_ngrams

bert/generation_utils.py:861–865  ·  view source on GitHub ↗
(hypo_idx)

Source from the content-addressed store, hash-verified

859 generated_ngram[prev_ngram_tuple] = generated_ngram.get(prev_ngram_tuple, []) + [ngram[-1]]
860
861 def _get_generated_ngrams(hypo_idx):
862 # Before decoding the next token, prevent decoding of ngrams that have already appeared
863 start_idx = cur_len + 1 - no_repeat_ngram_size
864 ngram_idx = tuple(prev_input_ids[hypo_idx, start_idx:cur_len].tolist())
865 return generated_ngrams[hypo_idx].get(ngram_idx, [])
866
867 banned_tokens = [_get_generated_ngrams(hypo_idx) for hypo_idx in range(num_hypos)]
868 return banned_tokens

Callers 1

calc_banned_ngram_tokensFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected