Aluode/PerceptionLabPortable
0
1# Copyright 2018 The HuggingFace Inc. team.2#3# Licensed under the Apache License, Version 2.0 (the "License");4# you may not use this file except in compliance with the License.5# You may obtain a copy of the License at6#7# http://www.apache.org/licenses/LICENSE-2.08#9# Unless required by applicable law or agreed to in writing, software10# distributed under the License is distributed on an "AS IS" BASIS,11# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.12# See the License for the specific language governing permissions and13# limitations under the License.14"""15Utilities to convert slow tokenizers in their fast tokenizers counterparts.16 17All the conversions are grouped here to gather SentencePiece dependencies outside of the fast tokenizers files and18allow to make our dependency on SentencePiece optional.19"""20 21import warnings22from functools import lru_cache23from typing import Optional24 25from packaging import version26from tokenizers import AddedToken, Regex, Tokenizer, decoders, normalizers, pre_tokenizers, processors27from tokenizers.models import BPE, Unigram, WordPiece28from tqdm import tqdm29 30from .utils import is_protobuf_available, is_sentencepiece_available, logging, requires_backends31from .utils.import_utils import PROTOBUF_IMPORT_ERROR32 33 34logger = logging.get_logger(__name__)35 36 37def import_protobuf(error_message=""):38 if is_sentencepiece_available():39 from sentencepiece import sentencepiece_model_pb240 41 return sentencepiece_model_pb242 if is_protobuf_available():43 import google.protobuf44 45 if version.parse(google.protobuf.__version__) < version.parse("4.0.0"):46 from transformers.utils import sentencepiece_model_pb247 else:48 from transformers.utils import sentencepiece_model_pb2_new as sentencepiece_model_pb249 return sentencepiece_model_pb250 else:51 raise ImportError(PROTOBUF_IMPORT_ERROR.format(error_message))52 53 54def _get_prepend_scheme(add_prefix_space: bool, original_tokenizer) -> str:55 if add_prefix_space:56 prepend_scheme = "always"57 if not getattr(original_tokenizer, "legacy", True):58 prepend_scheme = "first"59 else:60 prepend_scheme = "never"61 return prepend_scheme62 63 64def generate_merges(vocab, vocab_scores):65 reverse = vocab_scores is not None66 vocab_scores = dict(vocab_scores) if reverse else vocab67 68 merges = []69 for merge, piece_score in vocab_scores.items():70 local = []71 for index in range(1, len(merge)):72 piece_l, piece_r = merge[:index], merge[index:]73 if piece_l in vocab and piece_r in vocab:74 local.append((piece_l, piece_r, piece_score))75 local = sorted(local, key=lambda x: (vocab[x[0]], vocab[x[1]]))76 merges.extend(local)77 78 merges = sorted(merges, key=lambda val: (val[2], len(val[0]), len(val[1])), reverse=reverse)79 merges = [(val[0], val[1]) for val in merges]80 return merges81 82 83class SentencePieceExtractor:84 """85 Extractor implementation for SentencePiece trained models. https://github.com/google/sentencepiece86 """87 88 def __init__(self, model: str):89 requires_backends(self, "sentencepiece")90 from sentencepiece import SentencePieceProcessor91 92 self.sp = SentencePieceProcessor()93 self.sp.Load(model)94 95 def extract(self, vocab_scores=None) -> tuple[dict[str, int], list[tuple]]:96 """97 By default will return vocab and merges with respect to their order, by sending `vocab_scores` we're going to98 order the merges with respect to the piece scores instead.99 """100 sp = self.sp101 vocab = {sp.id_to_piece(index): index for index in range(sp.GetPieceSize())}102 103 merges = generate_merges(vocab, vocab_scores)104 105 return vocab, merges106 107 108class GemmaSentencePieceExtractor(SentencePieceExtractor):109 def extract(self, vocab_scores=None) -> tuple[dict[str, int], list[tuple]]:110 """111 By default will return vocab and merges with respect to their order, by sending `vocab_scores` we're going to112 order the merges with respect to the piece scores instead.113 """114 sp = self.sp115 vocab = {sp.id_to_piece(index): index for index in range(sp.GetPieceSize())}116 117 # If "\t" is missing in the vocab, we have to do this to support merges118 # "<0x09>" is the bytefallback for `\t`119 if "\t" not in vocab:120 vocab["\t"] = vocab.get("<0x09>")121 merges = generate_merges(vocab, vocab_scores)122 return vocab, merges123 124 125def check_number_comma(piece: str) -> bool:126 return len(piece) < 2 or piece[-1] != "," or not piece[-2].isdigit()127 128 129class Converter:130 def __init__(self, original_tokenizer):131 self.original_tokenizer = original_tokenizer132 133 def converted(self) -> Tokenizer:134 raise NotImplementedError()135 136 137class BertConverter(Converter):138 def converted(self) -> Tokenizer:139 vocab = self.original_tokenizer.vocab140 tokenizer = Tokenizer(WordPiece(vocab, unk_token=str(self.original_tokenizer.unk_token)))141 142 tokenize_chinese_chars = False143 strip_accents = False144 do_lower_case = False145 if hasattr(self.original_tokenizer, "basic_tokenizer"):146 tokenize_chinese_chars = self.original_tokenizer.basic_tokenizer.tokenize_chinese_chars147 strip_accents = self.original_tokenizer.basic_tokenizer.strip_accents148 do_lower_case = self.original_tokenizer.basic_tokenizer.do_lower_case149 150 tokenizer.normalizer = normalizers.BertNormalizer(151 clean_text=True,152 handle_chinese_chars=tokenize_chinese_chars,153 strip_accents=strip_accents,154 lowercase=do_lower_case,155 )156 tokenizer.pre_tokenizer = pre_tokenizers.BertPreTokenizer()157 158 cls = str(self.original_tokenizer.cls_token)159 sep = str(self.original_tokenizer.sep_token)160 cls_token_id = self.original_tokenizer.cls_token_id161 sep_token_id = self.original_tokenizer.sep_token_id162 163 tokenizer.post_processor = processors.TemplateProcessing(164 single=f"{cls}:0 $A:0 {sep}:0",165 pair=f"{cls}:0 $A:0 {sep}:0 $B:1 {sep}:1",166 special_tokens=[167 (cls, cls_token_id),168 (sep, sep_token_id),169 ],170 )171 tokenizer.decoder = decoders.WordPiece(prefix="##")172 173 return tokenizer174 175 176class SplinterConverter(Converter):177 def converted(self) -> Tokenizer:178 vocab = self.original_tokenizer.vocab179 tokenizer = Tokenizer(WordPiece(vocab, unk_token=str(self.original_tokenizer.unk_token)))180 181 tokenize_chinese_chars = False182 strip_accents = False183 do_lower_case = False184 if hasattr(self.original_tokenizer, "basic_tokenizer"):185 tokenize_chinese_chars = self.original_tokenizer.basic_tokenizer.tokenize_chinese_chars186 strip_accents = self.original_tokenizer.basic_tokenizer.strip_accents187 do_lower_case = self.original_tokenizer.basic_tokenizer.do_lower_case188 189 tokenizer.normalizer = normalizers.BertNormalizer(190 clean_text=True,191 handle_chinese_chars=tokenize_chinese_chars,192 strip_accents=strip_accents,193 lowercase=do_lower_case,194 )195 tokenizer.pre_tokenizer = pre_tokenizers.BertPreTokenizer()196 197 cls = str(self.original_tokenizer.cls_token)198 sep = str(self.original_tokenizer.sep_token)199 question = str(self.original_tokenizer.question_token)200 dot = "."201 cls_token_id = self.original_tokenizer.cls_token_id202 sep_token_id = self.original_tokenizer.sep_token_id203 question_token_id = self.original_tokenizer.question_token_id204 dot_token_id = self.original_tokenizer.convert_tokens_to_ids(".")205 206 if self.original_tokenizer.padding_side == "right":207 pair = f"{cls}:0 $A:0 {question} {dot} {sep}:0 $B:1 {sep}:1"208 else:209 pair = f"{cls}:0 $A:0 {sep}:0 $B:1 {question} {dot} {sep}:1"210 211 tokenizer.post_processor = processors.TemplateProcessing(212 single=f"{cls}:0 $A:0 {sep}:0",213 pair=pair,214 special_tokens=[215 (cls, cls_token_id),216 (sep, sep_token_id),217 (question, question_token_id),218 (dot, dot_token_id),219 ],220 )221 tokenizer.decoder = decoders.WordPiece(prefix="##")222 223 return tokenizer224 225 226class FunnelConverter(Converter):227 def converted(self) -> Tokenizer:228 vocab = self.original_tokenizer.vocab229 tokenizer = Tokenizer(WordPiece(vocab, unk_token=str(self.original_tokenizer.unk_token)))230 231 tokenize_chinese_chars = False232 strip_accents = False233 do_lower_case = False234 if hasattr(self.original_tokenizer, "basic_tokenizer"):235 tokenize_chinese_chars = self.original_tokenizer.basic_tokenizer.tokenize_chinese_chars236 strip_accents = self.original_tokenizer.basic_tokenizer.strip_accents237 do_lower_case = self.original_tokenizer.basic_tokenizer.do_lower_case238 239 tokenizer.normalizer = normalizers.BertNormalizer(240 clean_text=True,241 handle_chinese_chars=tokenize_chinese_chars,242 strip_accents=strip_accents,243 lowercase=do_lower_case,244 )245 tokenizer.pre_tokenizer = pre_tokenizers.BertPreTokenizer()246 247 cls = str(self.original_tokenizer.cls_token)248 sep = str(self.original_tokenizer.sep_token)249 cls_token_id = self.original_tokenizer.cls_token_id250 sep_token_id = self.original_tokenizer.sep_token_id251 252 tokenizer.post_processor = processors.TemplateProcessing(253 single=f"{cls}:2 $A:0 {sep}:0", # token_type_id is 2 for Funnel transformer254 pair=f"{cls}:2 $A:0 {sep}:0 $B:1 {sep}:1",255 special_tokens=[256 (cls, cls_token_id),257 (sep, sep_token_id),258 ],259 )260 tokenizer.decoder = decoders.WordPiece(prefix="##")261 262 return tokenizer263 264 265class MPNetConverter(Converter):266 def converted(self) -> Tokenizer:267 vocab = self.original_tokenizer.vocab268 tokenizer = Tokenizer(WordPiece(vocab, unk_token=str(self.original_tokenizer.unk_token)))269 270 tokenize_chinese_chars = False271 strip_accents = False272 do_lower_case = False273 if hasattr(self.original_tokenizer, "basic_tokenizer"):274 tokenize_chinese_chars = self.original_tokenizer.basic_tokenizer.tokenize_chinese_chars275 strip_accents = self.original_tokenizer.basic_tokenizer.strip_accents276 do_lower_case = self.original_tokenizer.basic_tokenizer.do_lower_case277 278 tokenizer.normalizer = normalizers.BertNormalizer(279 clean_text=True,280 handle_chinese_chars=tokenize_chinese_chars,281 strip_accents=strip_accents,282 lowercase=do_lower_case,283 )284 tokenizer.pre_tokenizer = pre_tokenizers.BertPreTokenizer()285 286 cls = str(self.original_tokenizer.cls_token)287 sep = str(self.original_tokenizer.sep_token)288 cls_token_id = self.original_tokenizer.cls_token_id289 sep_token_id = self.original_tokenizer.sep_token_id290 291 tokenizer.post_processor = processors.TemplateProcessing(292 single=f"{cls}:0 $A:0 {sep}:0",293 pair=f"{cls}:0 $A:0 {sep}:0 {sep}:0 $B:1 {sep}:1", # MPNet uses two [SEP] tokens294 special_tokens=[295 (cls, cls_token_id),296 (sep, sep_token_id),297 ],298 )299 tokenizer.decoder = decoders.WordPiece(prefix="##")300 301 return tokenizer302 303 304class OpenAIGPTConverter(Converter):305 def converted(self) -> Tokenizer:306 vocab = self.original_tokenizer.encoder307 merges = list(self.original_tokenizer.bpe_ranks.keys())308 unk_token = self.original_tokenizer.unk_token309 310 tokenizer = Tokenizer(311 BPE(312 vocab=vocab,313 merges=merges,314 dropout=None,315 unk_token=str(unk_token),316 end_of_word_suffix="</w>",317 fuse_unk=False,318 )319 )320 321 if tokenizer.token_to_id(str(unk_token)) is not None:322 tokenizer.add_special_tokens([str(unk_token)])323 324 tokenizer.normalizer = normalizers.BertNormalizer(lowercase=True)325 tokenizer.pre_tokenizer = pre_tokenizers.BertPreTokenizer()326 tokenizer.decoder = decoders.BPEDecoder(suffix="</w>")327 328 return tokenizer329 330 331class GPT2Converter(Converter):332 def converted(333 self, vocab: Optional[dict[str, int]] = None, merges: Optional[list[tuple[str, str]]] = None334 ) -> Tokenizer:335 if not vocab:336 vocab = self.original_tokenizer.encoder337 if not merges:338 merges = list(self.original_tokenizer.bpe_ranks)339 340 tokenizer = Tokenizer(341 BPE(342 vocab=vocab,343 merges=merges,344 dropout=None,345 continuing_subword_prefix="",346 end_of_word_suffix="",347 fuse_unk=False,348 )349 )350 351 add_prefix_space = getattr(self.original_tokenizer, "add_prefix_space", False)352 tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=add_prefix_space)353 tokenizer.decoder = decoders.ByteLevel()354 if getattr(self.original_tokenizer, "add_bos_token", False):355 bos = self.original_tokenizer.bos_token356 bos_token_id = self.original_tokenizer.bos_token_id357 tokenizer.post_processor = processors.TemplateProcessing(358 single=f"{bos}:0 $A:0",359 pair=f"{bos}:0 $A:0 $B:1",360 special_tokens=[361 (bos, bos_token_id),362 ],363 )364 else:365 # XXX trim_offsets=False actually means this post_processor doesn't366 # really do anything.367 tokenizer.post_processor = processors.ByteLevel(trim_offsets=False)368 return tokenizer369 370 371class HerbertConverter(Converter):372 def converted(self) -> Tokenizer:373 tokenizer_info_str = "#version:"374 token_suffix = "</w>"375 376 vocab = self.original_tokenizer.encoder377 merges = list(self.original_tokenizer.bpe_ranks.keys())378 if tokenizer_info_str in merges[0][0]:379 merges = merges[1:]380 381 tokenizer = Tokenizer(382 BPE(383 vocab,384 merges,385 dropout=None,386 unk_token=self.original_tokenizer.unk_token,387 end_of_word_suffix=token_suffix,388 )389 )390 391 tokenizer.normalizer = normalizers.BertNormalizer(lowercase=False, strip_accents=False)392 tokenizer.pre_tokenizer = pre_tokenizers.BertPreTokenizer()393 tokenizer.decoder = decoders.BPEDecoder(suffix=token_suffix)394 tokenizer.post_processor = processors.BertProcessing(395 sep=(self.original_tokenizer.sep_token, self.original_tokenizer.sep_token_id),396 cls=(self.original_tokenizer.cls_token, self.original_tokenizer.cls_token_id),397 )398 399 return tokenizer400 401 402class Qwen2Converter(Converter):403 def converted(404 self, vocab: Optional[dict[str, int]] = None, merges: Optional[list[tuple[str, str]]] = None405 ) -> Tokenizer:406 if not vocab:407 vocab = self.original_tokenizer.encoder408 if not merges:409 merges = list(self.original_tokenizer.bpe_ranks.keys())410 411 tokenizer = Tokenizer(412 BPE(413 vocab=vocab,414 merges=merges,415 dropout=None,416 unk_token=None,417 continuing_subword_prefix="",418 end_of_word_suffix="",419 fuse_unk=False,420 byte_fallback=False,421 )422 )423 424 tokenizer.normalizer = normalizers.NFC()425 426 tokenizer.pre_tokenizer = pre_tokenizers.Sequence(427 [428 pre_tokenizers.Split(429 Regex(430 r"""(?i:'s|'t|'re|'ve|'m|'ll|'d)|[^\r\n\p{L}\p{N}]?\p{L}+|\p{N}| ?[^\s\p{L}\p{N}]+[\r\n]*|\s*[\r\n]+|\s+(?!\S)|\s+"""431 ),432 behavior="isolated",433 invert=False,434 ),435 pre_tokenizers.ByteLevel(436 add_prefix_space=getattr(self.original_tokenizer, "add_prefix_space", False),437 use_regex=False,438 ),439 ]440 )441 442 tokenizer.decoder = decoders.ByteLevel()443 tokenizer.post_processor = processors.ByteLevel(trim_offsets=False)444 445 return tokenizer446 447 448class RobertaConverter(Converter):449 def converted(self) -> Tokenizer:450 ot = self.original_tokenizer451 vocab = ot.encoder452 merges = list(ot.bpe_ranks.keys())453 454 tokenizer = Tokenizer(455 BPE(456 vocab=vocab,457 merges=merges,458 dropout=None,459 continuing_subword_prefix="",460 end_of_word_suffix="",461 fuse_unk=False,462 )463 )464 465 tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=ot.add_prefix_space)466 tokenizer.decoder = decoders.ByteLevel()467 tokenizer.post_processor = processors.RobertaProcessing(468 sep=(ot.sep_token, ot.sep_token_id),469 cls=(ot.cls_token, ot.cls_token_id),470 add_prefix_space=ot.add_prefix_space,471 trim_offsets=True, # True by default on Roberta (historical)472 )473 474 return tokenizer475 476 477class RoFormerConverter(Converter):478 def converted(self) -> Tokenizer:479 from .models.roformer.tokenization_utils import JiebaPreTokenizer480 481 vocab = self.original_tokenizer.vocab482 tokenizer = Tokenizer(WordPiece(vocab, unk_token=str(self.original_tokenizer.unk_token)))483 484 strip_accents = False485 do_lower_case = False486 if hasattr(self.original_tokenizer, "basic_tokenizer"):487 strip_accents = self.original_tokenizer.basic_tokenizer.strip_accents488 do_lower_case = self.original_tokenizer.basic_tokenizer.do_lower_case489 490 tokenizer.normalizer = normalizers.BertNormalizer(491 clean_text=True,492 handle_chinese_chars=False,493 strip_accents=strip_accents,494 lowercase=do_lower_case,495 )496 tokenizer.pre_tokenizer = pre_tokenizers.PreTokenizer.custom(JiebaPreTokenizer(vocab))497 498 cls = str(self.original_tokenizer.cls_token)499 sep = str(self.original_tokenizer.sep_token)500 cls_token_id = self.original_tokenizer.cls_token_id501 sep_token_id = self.original_tokenizer.sep_token_id502 503 tokenizer.post_processor = processors.TemplateProcessing(504 single=f"{cls}:0 $A:0 {sep}:0",505 pair=f"{cls}:0 $A:0 {sep}:0 $B:1 {sep}:1",506 special_tokens=[507 (cls, cls_token_id),508 (sep, sep_token_id),509 ],510 )511 tokenizer.decoder = decoders.WordPiece(prefix="##")512 513 return tokenizer514 515 516class DebertaConverter(Converter):517 def converted(self) -> Tokenizer:518 ot = self.original_tokenizer519 vocab = ot.encoder520 merges = list(ot.bpe_ranks.keys())521 522 tokenizer = Tokenizer(523 BPE(524 vocab=vocab,525 merges=merges,526 dropout=None,527 continuing_subword_prefix="",528 end_of_word_suffix="",529 fuse_unk=False,530 )531 )532 533 tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=ot.add_prefix_space)534 tokenizer.decoder = decoders.ByteLevel()535 tokenizer.post_processor = processors.TemplateProcessing(536 single="[CLS]:0 $A:0 [SEP]:0",537 pair="[CLS]:0 $A:0 [SEP]:0 $B:1 [SEP]:1",538 special_tokens=[539 ("[CLS]", self.original_tokenizer.convert_tokens_to_ids("[CLS]")),540 ("[SEP]", self.original_tokenizer.convert_tokens_to_ids("[SEP]")),541 ],542 )543 544 return tokenizer545 546 547class SpmConverter(Converter):548 handle_byte_fallback = False549 SpmExtractor = SentencePieceExtractor550 special_tokens = {}551 552 def __init__(self, *args):553 requires_backends(self, "protobuf")554 555 super().__init__(*args)556 557 # from .utils import sentencepiece_model_pb2 as model_pb2558 model_pb2 = import_protobuf()559 560 m = model_pb2.ModelProto()561 with open(self.original_tokenizer.vocab_file, "rb") as f:562 m.ParseFromString(f.read())563 self.proto = m564 565 if self.proto.trainer_spec.byte_fallback and not self.handle_byte_fallback:566 warnings.warn(567 "The sentencepiece tokenizer that you are converting to a fast tokenizer uses the byte fallback option"568 " which is not implemented in the fast tokenizers. In practice this means that the fast version of the"569 " tokenizer can produce unknown tokens whereas the sentencepiece version would have converted these "570 "unknown tokens into a sequence of byte tokens matching the original piece of text."571 )572 573 def vocab(self, proto):574 return [(piece.piece, piece.score) for piece in proto.pieces]575 576 def unk_id(self, proto):577 return proto.trainer_spec.unk_id578 579 def tokenizer(self, proto):580 model_type = proto.trainer_spec.model_type581 vocab_scores = self.vocab(proto)582 583 if model_type == 1:584 tokenizer = Tokenizer(585 Unigram(586 vocab_scores,587 unk_id=self.unk_id(proto),588 byte_fallback=self.handle_byte_fallback,589 )590 )591 592 elif model_type == 2:593 _, merges = self.SpmExtractor(self.original_tokenizer.vocab_file).extract(vocab_scores)594 bpe_vocab = {word: i for i, (word, score) in enumerate(vocab_scores)}595 tokenizer = Tokenizer(596 BPE(597 bpe_vocab,598 merges,599 unk_token=proto.trainer_spec.unk_piece,600 fuse_unk=True,601 byte_fallback=self.handle_byte_fallback,602 dropout=None,603 )604 )605 606 else:607 raise Exception(608 "You're trying to run a `Unigram` model but you're file was trained with a different algorithm"609 )610 611 # control tokens are special612 # user defined symbols are not613 # both user and control tokens are AddedTokens614 # Add user defined symbols (type == 4) from sentencepiece (https://github.com/google/sentencepiece/blob/6225e08edb2577757163b3f5dbba4c0b670ef445/src/sentencepiece_model.proto#L299C29-L299C33)615 spm_added_tokens = [616 (id, p.piece, p.type == 3 or p.piece in self.special_tokens)617 for id, p in enumerate(proto.pieces)618 if p.type in [3, 4]619 ]620 tokenizer.add_tokens(621 [622 AddedToken(token, normalized=False, special=special)623 for id, token, special in sorted(spm_added_tokens, key=lambda x: x[0])624 ]625 )626 627 return tokenizer628 629 def normalizer(self, proto):630 precompiled_charsmap = proto.normalizer_spec.precompiled_charsmap631 _normalizers = [632 normalizers.Strip(left=False, right=True), # stripping is important633 normalizers.Replace(Regex(" {2,}"), "▁"),634 ]635 if not precompiled_charsmap:636 return normalizers.Sequence(_normalizers)637 else:638 return normalizers.Sequence([normalizers.Precompiled(precompiled_charsmap)] + _normalizers)639 640 def pre_tokenizer(self, replacement, add_prefix_space):641 prepend_scheme = _get_prepend_scheme(add_prefix_space, self.original_tokenizer)642 return pre_tokenizers.Metaspace(replacement=replacement, prepend_scheme=prepend_scheme)643 644 def post_processor(self):645 return None646 647 def decoder(self, replacement, add_prefix_space):648 prepend_scheme = _get_prepend_scheme(add_prefix_space, self.original_tokenizer)649 return decoders.Metaspace(replacement=replacement, prepend_scheme=prepend_scheme)650 651 def converted(self) -> Tokenizer:652 tokenizer = self.tokenizer(self.proto)653 654 # Tokenizer assemble655 normalizer = self.normalizer(self.proto)656 if normalizer is not None:657 tokenizer.normalizer = normalizer658 659 replacement = "▁"660 add_prefix_space = True661 if hasattr(self.original_tokenizer, "add_prefix_space"):662 add_prefix_space = self.original_tokenizer.add_prefix_space663 664 pre_tokenizer = self.pre_tokenizer(replacement, add_prefix_space)665 if pre_tokenizer is not None:666 tokenizer.pre_tokenizer = pre_tokenizer667 668 tokenizer.decoder = self.decoder(replacement, add_prefix_space)669 post_processor = self.post_processor()670 if post_processor:671 tokenizer.post_processor = post_processor672 673 return tokenizer674 675 676class AlbertConverter(SpmConverter):677 def vocab(self, proto):678 return [679 (piece.piece, piece.score) if check_number_comma(piece.piece) else (piece.piece, piece.score - 100)680 for piece in proto.pieces681 ]682 683 def normalizer(self, proto):684 list_normalizers = [685 normalizers.Replace("``", '"'),686 normalizers.Replace("''", '"'),687 ]688 if not self.original_tokenizer.keep_accents:689 list_normalizers.append(normalizers.NFKD())690 list_normalizers.append(normalizers.StripAccents())691 if self.original_tokenizer.do_lower_case:692 list_normalizers.append(normalizers.Lowercase())693 694 precompiled_charsmap = proto.normalizer_spec.precompiled_charsmap695 696 if precompiled_charsmap:697 list_normalizers.append(normalizers.Precompiled(precompiled_charsmap))698 699 list_normalizers.append(normalizers.Replace(Regex(" {2,}"), " "))700 return normalizers.Sequence(list_normalizers)701 702 def post_processor(self):703 return processors.TemplateProcessing(704 single="[CLS]:0 $A:0 [SEP]:0",705 pair="[CLS]:0 $A:0 [SEP]:0 $B:1 [SEP]:1",706 special_tokens=[707 ("[CLS]", self.original_tokenizer.convert_tokens_to_ids("[CLS]")),708 ("[SEP]", self.original_tokenizer.convert_tokens_to_ids("[SEP]")),709 ],710 )711 712 713class BarthezConverter(SpmConverter):714 def unk_id(self, proto):715 unk_id = 3716 return unk_id717 718 def post_processor(self):719 return processors.TemplateProcessing(720 single="<s> $A </s>",721 pair="<s> $A </s> </s> $B </s>",722 special_tokens=[723 ("<s>", self.original_tokenizer.convert_tokens_to_ids("<s>")),724 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),725 ],726 )727 728 729class CamembertConverter(SpmConverter):730 def vocab(self, proto):731 vocab = [732 ("<s>NOTUSED", 0.0),733 ("<pad>", 0.0),734 ("</s>NOTUSED", 0.0),735 ("<unk>", 0.0),736 ("<unk>NOTUSED", -100),737 ]738 # We down-grade the original SentencePiece by -100 to avoid using it and use our added token instead739 vocab += [(piece.piece, piece.score) for piece in proto.pieces[1:]]740 vocab += [("<mask>", 0.0)]741 return vocab742 743 def unk_id(self, proto):744 # See vocab unk position745 return 3746 747 def post_processor(self):748 return processors.TemplateProcessing(749 single="<s> $A </s>",750 pair="<s> $A </s> </s> $B </s>",751 special_tokens=[752 ("<s>", self.original_tokenizer.convert_tokens_to_ids("<s>")),753 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),754 ],755 )756 757 758class DebertaV2Converter(SpmConverter):759 def pre_tokenizer(self, replacement, add_prefix_space):760 list_pretokenizers = []761 if self.original_tokenizer.split_by_punct:762 list_pretokenizers.append(pre_tokenizers.Punctuation(behavior="isolated"))763 prepend_scheme = _get_prepend_scheme(add_prefix_space, self.original_tokenizer)764 list_pretokenizers.append(pre_tokenizers.Metaspace(replacement=replacement, prepend_scheme=prepend_scheme))765 return pre_tokenizers.Sequence(list_pretokenizers)766 767 def normalizer(self, proto):768 list_normalizers = []769 if self.original_tokenizer.do_lower_case:770 list_normalizers.append(normalizers.Lowercase())771 list_normalizers.append(normalizers.Strip())772 773 precompiled_charsmap = proto.normalizer_spec.precompiled_charsmap774 if precompiled_charsmap:775 list_normalizers.append(normalizers.Precompiled(precompiled_charsmap))776 list_normalizers.append(normalizers.Replace(Regex(" {2,}"), " "))777 778 return normalizers.Sequence(list_normalizers)779 780 def post_processor(self):781 return processors.TemplateProcessing(782 single="[CLS]:0 $A:0 [SEP]:0",783 pair="[CLS]:0 $A:0 [SEP]:0 $B:1 [SEP]:1",784 special_tokens=[785 ("[CLS]", self.original_tokenizer.convert_tokens_to_ids("[CLS]")),786 ("[SEP]", self.original_tokenizer.convert_tokens_to_ids("[SEP]")),787 ],788 )789 790 791class MBartConverter(SpmConverter):792 def vocab(self, proto):793 vocab = [794 ("<s>", 0.0),795 ("<pad>", 0.0),796 ("</s>", 0.0),797 ("<unk>", 0.0),798 ]799 vocab += [(piece.piece, piece.score) for piece in proto.pieces[3:]]800 vocab += [801 ("ar_AR", 0.0),802 ("cs_CZ", 0.0),803 ("de_DE", 0.0),804 ("en_XX", 0.0),805 ("es_XX", 0.0),806 ("et_EE", 0.0),807 ("fi_FI", 0.0),808 ("fr_XX", 0.0),809 ("gu_IN", 0.0),810 ("hi_IN", 0.0),811 ("it_IT", 0.0),812 ("ja_XX", 0.0),813 ("kk_KZ", 0.0),814 ("ko_KR", 0.0),815 ("lt_LT", 0.0),816 ("lv_LV", 0.0),817 ("my_MM", 0.0),818 ("ne_NP", 0.0),819 ("nl_XX", 0.0),820 ("ro_RO", 0.0),821 ("ru_RU", 0.0),822 ("si_LK", 0.0),823 ("tr_TR", 0.0),824 ("vi_VN", 0.0),825 ("zh_CN", 0.0),826 ]827 vocab += [("<mask>", 0.0)]828 return vocab829 830 def unk_id(self, proto):831 return 3832 833 def post_processor(self):834 return processors.TemplateProcessing(835 single="$A </s> en_XX",836 pair="$A $B </s> en_XX",837 special_tokens=[838 ("en_XX", self.original_tokenizer.convert_tokens_to_ids("en_XX")),839 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),840 ],841 )842 843 844class MBart50Converter(SpmConverter):845 def vocab(self, proto):846 vocab = [847 ("<s>", 0.0),848 ("<pad>", 0.0),849 ("</s>", 0.0),850 ("<unk>", 0.0),851 ]852 vocab += [(piece.piece, piece.score) for piece in proto.pieces[3:]]853 vocab += [("ar_AR", 0.0), ("cs_CZ", 0.0), ("de_DE", 0.0), ("en_XX", 0.0), ("es_XX", 0.0), ("et_EE", 0.0), ("fi_FI", 0.0), ("fr_XX", 0.0), ("gu_IN", 0.0), ("hi_IN", 0.0), ("it_IT", 0.0), ("ja_XX", 0.0), ("kk_KZ", 0.0), ("ko_KR", 0.0), ("lt_LT", 0.0), ("lv_LV", 0.0), ("my_MM", 0.0), ("ne_NP", 0.0), ("nl_XX", 0.0), ("ro_RO", 0.0), ("ru_RU", 0.0), ("si_LK", 0.0), ("tr_TR", 0.0), ("vi_VN", 0.0), ("zh_CN", 0.0), ("af_ZA", 0.0), ("az_AZ", 0.0), ("bn_IN", 0.0), ("fa_IR", 0.0), ("he_IL", 0.0), ("hr_HR", 0.0), ("id_ID", 0.0), ("ka_GE", 0.0), ("km_KH", 0.0), ("mk_MK", 0.0), ("ml_IN", 0.0), ("mn_MN", 0.0), ("mr_IN", 0.0), ("pl_PL", 0.0), ("ps_AF", 0.0), ("pt_XX", 0.0), ("sv_SE", 0.0), ("sw_KE", 0.0), ("ta_IN", 0.0), ("te_IN", 0.0), ("th_TH", 0.0), ("tl_XX", 0.0), ("uk_UA", 0.0), ("ur_PK", 0.0), ("xh_ZA", 0.0), ("gl_ES", 0.0), ("sl_SI", 0.0)] # fmt: skip854 vocab += [("<mask>", 0.0)]855 return vocab856 857 def unk_id(self, proto):858 return 3859 860 def post_processor(self):861 return processors.TemplateProcessing(862 single="en_XX $A </s>",863 pair="en_XX $A $B </s>",864 special_tokens=[865 ("en_XX", self.original_tokenizer.convert_tokens_to_ids("en_XX")),866 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),867 ],868 )869 870 871class NllbConverter(SpmConverter):872 def vocab(self, proto):873 vocab = [874 ("<s>", 0.0),875 ("<pad>", 0.0),876 ("</s>", 0.0),877 ("<unk>", 0.0),878 ]879 vocab += [(piece.piece, piece.score) for piece in proto.pieces[3:]]880 return vocab881 882 def unk_id(self, proto):883 return 3884 885 def post_processor(self):886 return processors.TemplateProcessing(887 single="eng_Latn $A </s>",888 pair="eng_Latn $A $B </s>",889 special_tokens=[890 ("eng_Latn", self.original_tokenizer.convert_tokens_to_ids("eng_Latn")),891 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),892 ],893 )894 895 896class SeamlessM4TConverter(SpmConverter):897 def vocab(self, proto):898 vocab = [899 ("<pad>", 0.0),900 ("<unk>", 0.0),901 ("<s>", 0.0),902 ("</s>", 0.0),903 ]904 vocab += [(piece.piece, piece.score) for piece in proto.pieces[3:]]905 return vocab906 907 def unk_id(self, proto):908 return self.original_tokenizer.unk_token_id909 910 def post_processor(self):911 return processors.TemplateProcessing(912 single="__eng__ $A </s>",913 pair="__eng__ $A $B </s>",914 special_tokens=[915 ("__eng__", self.original_tokenizer.convert_tokens_to_ids("__eng__")),916 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),917 ],918 )919 920 921class XLMRobertaConverter(SpmConverter):922 def vocab(self, proto):923 vocab = [924 ("<s>", 0.0),925 ("<pad>", 0.0),926 ("</s>", 0.0),927 ("<unk>", 0.0),928 ]929 vocab += [(piece.piece, piece.score) for piece in proto.pieces[3:]]930 vocab += [("<mask>", 0.0)]931 return vocab932 933 def unk_id(self, proto):934 unk_id = 3935 return unk_id936 937 def post_processor(self):938 return processors.TemplateProcessing(939 single="<s> $A </s>",940 pair="<s> $A </s> </s> $B </s>",941 special_tokens=[942 ("<s>", self.original_tokenizer.convert_tokens_to_ids("<s>")),943 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),944 ],945 )946 947 948class XLNetConverter(SpmConverter):949 def vocab(self, proto):950 return [951 (piece.piece, piece.score) if check_number_comma(piece.piece) else (piece.piece, piece.score - 100)952 for piece in proto.pieces953 ]954 955 def normalizer(self, proto):956 list_normalizers = [957 normalizers.Replace("``", '"'),958 normalizers.Replace("''", '"'),959 ]960 if not self.original_tokenizer.keep_accents:961 list_normalizers.append(normalizers.NFKD())962 list_normalizers.append(normalizers.StripAccents())963 if self.original_tokenizer.do_lower_case:964 list_normalizers.append(normalizers.Lowercase())965 966 precompiled_charsmap = proto.normalizer_spec.precompiled_charsmap967 968 if precompiled_charsmap:969 list_normalizers.append(normalizers.Precompiled(precompiled_charsmap))970 971 list_normalizers.append(normalizers.Replace(Regex(" {2,}"), " "))972 return normalizers.Sequence(list_normalizers)973 974 def post_processor(self):975 return processors.TemplateProcessing(976 single="$A:0 <sep>:0 <cls>:2",977 pair="$A:0 <sep>:0 $B:1 <sep>:1 <cls>:2",978 special_tokens=[979 ("<sep>", self.original_tokenizer.convert_tokens_to_ids("<sep>")),980 ("<cls>", self.original_tokenizer.convert_tokens_to_ids("<cls>")),981 ],982 )983 984 985class ReformerConverter(SpmConverter):986 pass987 988 989class RemBertConverter(SpmConverter):990 # Inspired from AlbertConverter991 def normalizer(self, proto):992 list_normalizers = [993 normalizers.Replace("``", '"'),994 normalizers.Replace("''", '"'),995 normalizers.Replace(Regex(" {2,}"), " "),996 ]997 if not self.original_tokenizer.keep_accents:998 list_normalizers.append(normalizers.NFKD())999 list_normalizers.append(normalizers.StripAccents())1000 if self.original_tokenizer.do_lower_case:1001 list_normalizers.append(normalizers.Lowercase())1002 1003 precompiled_charsmap = proto.normalizer_spec.precompiled_charsmap1004 1005 if precompiled_charsmap:1006 list_normalizers.append(normalizers.Precompiled(precompiled_charsmap))1007 1008 return normalizers.Sequence(list_normalizers)1009 1010 def post_processor(self):1011 return processors.TemplateProcessing(1012 single="[CLS]:0 $A:0 [SEP]:0",1013 pair="[CLS]:0 $A:0 [SEP]:0 $B:1 [SEP]:1",1014 special_tokens=[1015 ("[CLS]", self.original_tokenizer.convert_tokens_to_ids("[CLS]")),1016 ("[SEP]", self.original_tokenizer.convert_tokens_to_ids("[SEP]")),1017 ],1018 )1019 1020 1021class BertGenerationConverter(SpmConverter):1022 pass1023 1024 1025class PegasusConverter(SpmConverter):1026 def vocab(self, proto):1027 vocab = [1028 (self.original_tokenizer.pad_token, 0.0),1029 (self.original_tokenizer.eos_token, 0.0),1030 ]1031 1032 if self.original_tokenizer.mask_token_sent is not None:1033 vocab += [(self.original_tokenizer.mask_token_sent, 0.0)]1034 1035 if (1036 self.original_tokenizer.mask_token is not None1037 and self.original_tokenizer.mask_token_id < self.original_tokenizer.offset1038 ):1039 vocab += [(self.original_tokenizer.mask_token, 0.0)]1040 1041 vocab += [(f"<unk_{i}>", -100.0) for i in range(2, self.original_tokenizer.offset)]1042 vocab += [(piece.piece, piece.score) for piece in proto.pieces[2:]]1043 return vocab1044 1045 def unk_id(self, proto):1046 return proto.trainer_spec.unk_id + self.original_tokenizer.offset1047 1048 def pre_tokenizer(self, replacement, add_prefix_space):1049 prepend_scheme = _get_prepend_scheme(add_prefix_space, self.original_tokenizer)1050 return pre_tokenizers.Sequence(1051 [1052 pre_tokenizers.WhitespaceSplit(),1053 pre_tokenizers.Metaspace(replacement=replacement, prepend_scheme=prepend_scheme),1054 ]1055 )1056 1057 def post_processor(self):1058 eos = self.original_tokenizer.eos_token1059 special_tokens = [1060 (eos, self.original_tokenizer.eos_token_id),1061 ]1062 return processors.TemplateProcessing(single=["$A", eos], pair=["$A", "$B", eos], special_tokens=special_tokens)1063 1064 1065class T5Converter(SpmConverter):1066 def vocab(self, proto):1067 num_extra_ids = self.original_tokenizer._extra_ids1068 vocab = [(piece.piece, piece.score) for piece in proto.pieces]1069 vocab += [(f"<extra_id_{i}>", 0.0) for i in range(num_extra_ids - 1, -1, -1)]1070 return vocab1071 1072 def post_processor(self):1073 return processors.TemplateProcessing(1074 single=["$A", "</s>"],1075 pair=["$A", "</s>", "$B", "</s>"],1076 special_tokens=[1077 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),1078 ],1079 )1080 1081 1082class UdopConverter(SpmConverter):1083 def post_processor(self):1084 return processors.TemplateProcessing(1085 single=["$A", "</s>"],1086 pair=["$A", "</s>", "$B", "</s>"],1087 special_tokens=[1088 ("</s>", self.original_tokenizer.convert_tokens_to_ids("</s>")),1089 ],1090 )1091 1092 1093class WhisperConverter(Converter):1094 def converted(self) -> Tokenizer:1095 vocab = self.original_tokenizer.encoder1096 merges = list(self.original_tokenizer.bpe_ranks.keys())1097 1098 tokenizer = Tokenizer(1099 BPE(1100 vocab=vocab,1101 merges=merges,1102 dropout=None,1103 continuing_subword_prefix="",1104 end_of_word_suffix="",1105 fuse_unk=False,1106 )1107 )1108 1109 tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(add_prefix_space=self.original_tokenizer.add_prefix_space)1110 tokenizer.decoder = decoders.ByteLevel()1111 1112 prefix_token_ids = self.original_tokenizer.prefix_tokens1113 prefixes = self.original_tokenizer.convert_ids_to_tokens(prefix_token_ids)1114 eos = self.original_tokenizer.eos_token1115 eos_token_id = self.original_tokenizer.eos_token_id1116 prefix_template = " ".join([f"{token}:0" for token in prefixes])1117 tokenizer.post_processor = processors.TemplateProcessing(1118 single=f"{prefix_template} $A:0 {eos}:0",1119 pair=f"{prefix_template} $A:0 $B:1 {eos}:1",1120 special_tokens=[1121 (eos, eos_token_id),1122 *zip(prefixes, prefix_token_ids),1123 ],1124 )1125 1126 return tokenizer1127 1128 1129class BigBirdConverter(SpmConverter):1130 def post_processor(self):1131 return processors.TemplateProcessing(1132 single="[CLS]:0 $A:0 [SEP]:0",1133 pair="[CLS]:0 $A:0 [SEP]:0 $B:1 [SEP]:1",1134 special_tokens=[1135 ("[CLS]", self.original_tokenizer.convert_tokens_to_ids("[CLS]")),1136 ("[SEP]", self.original_tokenizer.convert_tokens_to_ids("[SEP]")),1137 ],1138 )1139 1140 1141class CLIPConverter(Converter):1142 def converted(self) -> Tokenizer:1143 vocab = self.original_tokenizer.encoder1144 merges = list(self.original_tokenizer.bpe_ranks.keys())1145 unk_token = self.original_tokenizer.unk_token1146 1147 tokenizer = Tokenizer(1148 BPE(1149 vocab=vocab,1150 merges=merges,1151 dropout=None,1152 continuing_subword_prefix="",1153 end_of_word_suffix="</w>",1154 fuse_unk=False,1155 unk_token=str(unk_token),1156 )1157 )1158 1159 tokenizer.normalizer = normalizers.Sequence(1160 [normalizers.NFC(), normalizers.Replace(Regex(r"\s+"), " "), normalizers.Lowercase()]1161 )1162 tokenizer.pre_tokenizer = pre_tokenizers.Sequence(1163 [1164 pre_tokenizers.Split(1165 Regex(r"""'s|'t|'re|'ve|'m|'ll|'d|[\p{L}]+|[\p{N}]|[^\s\p{L}\p{N}]+"""),1166 behavior="removed",1167 invert=True,1168 ),1169 pre_tokenizers.ByteLevel(add_prefix_space=False),1170 ]1171 )1172 tokenizer.decoder = decoders.ByteLevel()1173 1174 # Hack to have a ByteLevel and TemplaceProcessor1175 tokenizer.post_processor = processors.RobertaProcessing(1176 sep=(self.original_tokenizer.eos_token, self.original_tokenizer.eos_token_id),1177 cls=(self.original_tokenizer.bos_token, self.original_tokenizer.bos_token_id),1178 add_prefix_space=False,1179 trim_offsets=False,1180 )1181 return tokenizer1182 1183 1184class LayoutLMv2Converter(Converter):1185 def converted(self) -> Tokenizer:1186 vocab = self.original_tokenizer.vocab1187 tokenizer = Tokenizer(WordPiece(vocab, unk_token=str(self.original_tokenizer.unk_token)))1188 1189 tokenize_chinese_chars = False1190 strip_accents = False1191 do_lower_case = True1192 if hasattr(self.original_tokenizer, "basic_tokenizer"):1193 tokenize_chinese_chars = self.original_tokenizer.basic_tokenizer.tokenize_chinese_chars1194 strip_accents = self.original_tokenizer.basic_tokenizer.strip_accents1195 do_lower_case = self.original_tokenizer.basic_tokenizer.do_lower_case1196 1197 tokenizer.normalizer = normalizers.BertNormalizer(1198 clean_text=True,1199 handle_chinese_chars=tokenize_chinese_chars,1200 strip_accents=strip_accents,