Loads a pretrained WordPiece tokenizer from `cache_dir` for tokenization in BERT training. Default to bert-large-uncased tokenizer.
| 782 | |
| 783 | |
| 784 | class BertWordPieceTokenizer(Tokenizer): |
| 785 | """ |
| 786 | Loads a pretrained WordPiece tokenizer from `cache_dir` for tokenization |
| 787 | in BERT training. Default to bert-large-uncased tokenizer. |
| 788 | """ |
| 789 | |
| 790 | def __init__(self, tokenizer_model_type=None, cache_dir=None, add_block_symbols=False, add_sentinel_token=0, |
| 791 | add_task_mask=False, add_decoder_mask=False, **kwargs): |
| 792 | # default to bert-large-uncased tokenizer |
| 793 | if tokenizer_model_type not in PRETRAINED_VOCAB_ARCHIVE_MAP: |
| 794 | tokenizer_model_type = 'bert-large-uncased' |
| 795 | if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: |
| 796 | print('loading BertWordPieceTokenizer (', tokenizer_model_type, ') from cache_dir ', cache_dir) |
| 797 | do_lower_case = not ('-cased' in tokenizer_model_type or 'chinese' in tokenizer_model_type) |
| 798 | self.text_tokenizer = BertTokenizer.from_pretrained(tokenizer_model_type, do_lower_case=do_lower_case, |
| 799 | cache_dir=cache_dir) |
| 800 | if not torch.distributed.is_initialized() or torch.distributed.get_rank() == 0: |
| 801 | print('loaded', tokenizer_model_type) |
| 802 | # disable max len warnings by increasing max len |
| 803 | self.text_tokenizer.max_len = int(1e12) |
| 804 | |
| 805 | # set command tokens from wordpiece tokenizer values |
| 806 | self.num_command_tokens = 6 |
| 807 | self.num_tokens = len(self.text_tokenizer.vocab) |
| 808 | self.num_text_tokens = self.num_tokens - 5 |
| 809 | self.num_type_tokens = 2 |
| 810 | |
| 811 | self._command_tokens = [ |
| 812 | CommandToken('pad', '[PAD]', self.text_tokenizer.vocab['[PAD]']), |
| 813 | CommandToken('ENC', '[CLS]', self.text_tokenizer.vocab['[CLS]']), |
| 814 | CommandToken('MASK', '[MASK]', self.text_tokenizer.vocab['[MASK]']), |
| 815 | CommandToken('unk', '[UNK]', self.text_tokenizer.vocab['[UNK]']), |
| 816 | CommandToken('sep', '[SEP]', self.text_tokenizer.vocab['[SEP]']), |
| 817 | CommandToken('eos', '[PAD]', self.text_tokenizer.vocab['[PAD]']), |
| 818 | ] |
| 819 | if add_block_symbols: |
| 820 | self._command_tokens.extend([ |
| 821 | CommandToken('sop', '<|startofpiece|>', self.num_tokens), |
| 822 | CommandToken('eop', '<|endofpiece|>', self.num_tokens + 1) |
| 823 | ]) |
| 824 | self.num_tokens += 2 |
| 825 | self.num_command_tokens += 2 |
| 826 | if add_task_mask: |
| 827 | self._command_tokens.extend([ |
| 828 | CommandToken('gMASK', '[gMASK]', self.num_tokens), |
| 829 | CommandToken('sMASK', '[sMASK]', self.num_tokens + 1) |
| 830 | ]) |
| 831 | self.num_tokens += 2 |
| 832 | self.num_command_tokens += 2 |
| 833 | if add_decoder_mask: |
| 834 | self._command_tokens.extend([ |
| 835 | CommandToken('dBLOCK', '[dBLOCK]', self.num_tokens) |
| 836 | ]) |
| 837 | self.num_tokens += 1 |
| 838 | self.num_command_tokens += 1 |
| 839 | if add_sentinel_token > 0: |
| 840 | for i in range(1, add_sentinel_token): |
| 841 | self._command_tokens.extend([CommandToken(f'MASK{i}', f'[MASK{i}]', self.num_tokens), |