Support bad words
(self, bad_words, bad_words_token_ids)
| 761 | return stop_seqs, stop_seqs_len |
| 762 | |
| 763 | def update_bad_words(self, bad_words, bad_words_token_ids): |
| 764 | """Support bad words""" |
| 765 | |
| 766 | token_ids = bad_words_token_ids |
| 767 | |
| 768 | if token_ids is None: |
| 769 | token_ids = [] |
| 770 | for bad_word in bad_words: |
| 771 | # To prohibit words both at the beginning |
| 772 | # and in the middle of text |
| 773 | # (related to add_prefix_space tokenizer parameter) |
| 774 | for add_prefix_space in [False, True]: |
| 775 | prefix = " " if add_prefix_space else "" |
| 776 | prompt = prefix + bad_word.lstrip() |
| 777 | prompt_token_ids = self.tokenizer.convert_tokens_to_ids(self.tokenizer.tokenize(prompt)) |
| 778 | |
| 779 | if len(prompt_token_ids) != 1: |
| 780 | if not add_prefix_space: |
| 781 | data_processor_logger.warning( |
| 782 | f"Skip bad_words: <{prompt}>." |
| 783 | f"Bad words should be a single token." |
| 784 | f"Got tokens: {prompt_token_ids}." |
| 785 | ) |
| 786 | continue |
| 787 | |
| 788 | if prompt_token_ids[0] > self.tokenizer.vocab_size: |
| 789 | if not add_prefix_space: |
| 790 | data_processor_logger.warning( |
| 791 | f"Skip bad_words: <{prompt}>." |
| 792 | f"All token id values should be satisfying:" |
| 793 | f" 0 <= token_id < {self.tokenizer.vocab_size}." |
| 794 | f"Got token: {prompt_token_ids}." |
| 795 | ) |
| 796 | continue |
| 797 | |
| 798 | if prompt_token_ids not in token_ids: |
| 799 | token_ids.extend(prompt_token_ids) |
| 800 | return token_ids |