(hypo_idx)
| 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 |
no outgoing calls
no test coverage detected