codekingpro/portable-devtools
114k
1from typing import Dict, Iterator, List, Optional, Union
2
3from tokenizers import AddedToken, Tokenizer, decoders, trainers
4from tokenizers.models import WordPiece
5from tokenizers.normalizers import BertNormalizer
6from tokenizers.pre_tokenizers import BertPreTokenizer
7from tokenizers.processors import BertProcessing
8
9from .base_tokenizer import BaseTokenizer
10
11
12class BertWordPieceTokenizer(BaseTokenizer):
13 """Bert WordPiece Tokenizer"""
14
15 def __init__(
16 self,
17 vocab: Optional[Union[str, Dict[str, int]]] = None,
18 unk_token: Union[str, AddedToken] = "[UNK]",
19 sep_token: Union[str, AddedToken] = "[SEP]",
20 cls_token: Union[str, AddedToken] = "[CLS]",
21 pad_token: Union[str, AddedToken] = "[PAD]",
22 mask_token: Union[str, AddedToken] = "[MASK]",
23 clean_text: bool = True,
24 handle_chinese_chars: bool = True,
25 strip_accents: Optional[bool] = None,
26 lowercase: bool = True,
27 wordpieces_prefix: str = "##",
28 ):
29 if vocab is not None:
30 tokenizer = Tokenizer(WordPiece(vocab, unk_token=str(unk_token)))
31 else:
32 tokenizer = Tokenizer(WordPiece(unk_token=str(unk_token)))
33
34 # Let the tokenizer know about special tokens if they are part of the vocab
35 if tokenizer.token_to_id(str(unk_token)) is not None:
36 tokenizer.add_special_tokens([str(unk_token)])
37 if tokenizer.token_to_id(str(sep_token)) is not None:
38 tokenizer.add_special_tokens([str(sep_token)])
39 if tokenizer.token_to_id(str(cls_token)) is not None:
40 tokenizer.add_special_tokens([str(cls_token)])
41 if tokenizer.token_to_id(str(pad_token)) is not None:
42 tokenizer.add_special_tokens([str(pad_token)])
43 if tokenizer.token_to_id(str(mask_token)) is not None:
44 tokenizer.add_special_tokens([str(mask_token)])
45
46 tokenizer.normalizer = BertNormalizer(
47 clean_text=clean_text,
48 handle_chinese_chars=handle_chinese_chars,
49 strip_accents=strip_accents,
50 lowercase=lowercase,
51 )
52 tokenizer.pre_tokenizer = BertPreTokenizer()
53
54 if vocab is not None:
55 sep_token_id = tokenizer.token_to_id(str(sep_token))
56 if sep_token_id is None:
57 raise TypeError("sep_token not found in the vocabulary")
58 cls_token_id = tokenizer.token_to_id(str(cls_token))
59 if cls_token_id is None:
60 raise TypeError("cls_token not found in the vocabulary")
61
62 tokenizer.post_processor = BertProcessing((str(sep_token), sep_token_id), (str(cls_token), cls_token_id))
63 tokenizer.decoder = decoders.WordPiece(prefix=wordpieces_prefix)
64
65 parameters = {
66 "model": "BertWordPiece",
67 "unk_token": unk_token,
68 "sep_token": sep_token,
69 "cls_token": cls_token,
70 "pad_token": pad_token,
71 "mask_token": mask_token,
72 "clean_text": clean_text,
73 "handle_chinese_chars": handle_chinese_chars,
74 "strip_accents": strip_accents,
75 "lowercase": lowercase,
76 "wordpieces_prefix": wordpieces_prefix,
77 }
78
79 super().__init__(tokenizer, parameters)
80
81 @staticmethod
82 def from_file(vocab: str, **kwargs):
83 vocab = WordPiece.read_file(vocab)
84 return BertWordPieceTokenizer(vocab, **kwargs)
85
86 def train(
87 self,
88 files: Union[str, List[str]],
89 vocab_size: int = 30000,
90 min_frequency: int = 2,
91 limit_alphabet: int = 1000,
92 initial_alphabet: List[str] = [],
93 special_tokens: List[Union[str, AddedToken]] = [
94 "[PAD]",
95 "[UNK]",
96 "[CLS]",
97 "[SEP]",
98 "[MASK]",
99 ],
100 show_progress: bool = True,
101 wordpieces_prefix: str = "##",
102 ):
103 """Train the model using the given files"""
104
105 trainer = trainers.WordPieceTrainer(
106 vocab_size=vocab_size,
107 min_frequency=min_frequency,
108 limit_alphabet=limit_alphabet,
109 initial_alphabet=initial_alphabet,
110 special_tokens=special_tokens,
111 show_progress=show_progress,
112 continuing_subword_prefix=wordpieces_prefix,
113 )
114 if isinstance(files, str):
115 files = [files]
116 self._tokenizer.train(files, trainer=trainer)
117
118 def train_from_iterator(
119 self,
120 iterator: Union[Iterator[str], Iterator[Iterator[str]]],
121 vocab_size: int = 30000,
122 min_frequency: int = 2,
123 limit_alphabet: int = 1000,
124 initial_alphabet: List[str] = [],
125 special_tokens: List[Union[str, AddedToken]] = [
126 "[PAD]",
127 "[UNK]",
128 "[CLS]",
129 "[SEP]",
130 "[MASK]",
131 ],
132 show_progress: bool = True,
133 wordpieces_prefix: str = "##",
134 length: Optional[int] = None,
135 ):
136 """Train the model using the given iterator"""
137
138 trainer = trainers.WordPieceTrainer(
139 vocab_size=vocab_size,
140 min_frequency=min_frequency,
141 limit_alphabet=limit_alphabet,
142 initial_alphabet=initial_alphabet,
143 special_tokens=special_tokens,
144 show_progress=show_progress,
145 continuing_subword_prefix=wordpieces_prefix,
146 )
147 self._tokenizer.train_from_iterator(
148 iterator,
149 trainer=trainer,
150 length=length,
151 )
152 