timorobrecht/10k_initials_gpu1
02
1"""SentencePiece tokenizer wrapper for pinyin-code Transformers models."""
2
3from __future__ import annotations
4
5import shutil
6from pathlib import Path
7
8import sentencepiece as spm
9from transformers import PreTrainedTokenizer
10
11
12class PinyinCodeTokenizer(PreTrainedTokenizer):
13 """Slow tokenizer that preserves the existing SentencePiece model."""
14
15 vocab_files_names = {"vocab_file": "tokenizer.model"}
16 model_input_names = ["input_ids", "attention_mask"]
17
18 def __init__(
19 self,
20 vocab_file: str,
21 add_bos_token: bool = False,
22 add_eos_token: bool = False,
23 **kwargs,
24 ) -> None:
25 self.vocab_file = vocab_file
26 self.sp_model = spm.SentencePieceProcessor(model_file=vocab_file)
27 self.add_bos_token = add_bos_token
28 self.add_eos_token = add_eos_token
29
30 kwargs.setdefault("unk_token", self._piece_or_none(self.sp_model.unk_id()))
31 kwargs.setdefault("bos_token", self._piece_or_none(self.sp_model.bos_id()))
32 kwargs.setdefault("eos_token", self._piece_or_none(self.sp_model.eos_id()))
33 kwargs.setdefault("pad_token", self._piece_or_none(self.sp_model.pad_id()))
34 super().__init__(**kwargs)
35
36 def _piece_or_none(self, token_id: int) -> str | None:
37 if token_id is None or token_id < 0:
38 return None
39 return self.sp_model.id_to_piece(token_id)
40
41 @property
42 def vocab_size(self) -> int:
43 return self.sp_model.get_piece_size()
44
45 def get_vocab(self) -> dict[str, int]:
46 vocab = {self.sp_model.id_to_piece(i): i for i in range(self.vocab_size)}
47 vocab.update(self.added_tokens_encoder)
48 return vocab
49
50 def _tokenize(self, text: str) -> list[str]:
51 return self.sp_model.encode(text, out_type=str)
52
53 def _convert_token_to_id(self, token: str) -> int:
54 return self.sp_model.piece_to_id(token)
55
56 def _convert_id_to_token(self, index: int) -> str:
57 return self.sp_model.id_to_piece(index)
58
59 def convert_tokens_to_string(self, tokens: list[str]) -> str:
60 return self.sp_model.decode(tokens)
61
62 def build_inputs_with_special_tokens(
63 self,
64 token_ids_0: list[int],
65 token_ids_1: list[int] | None = None,
66 ) -> list[int]:
67 output = list(token_ids_0)
68 if self.add_bos_token and self.bos_token_id is not None:
69 output = [self.bos_token_id] + output
70 if self.add_eos_token and self.eos_token_id is not None:
71 output = output + [self.eos_token_id]
72 if token_ids_1 is not None:
73 output += list(token_ids_1)
74 if self.add_eos_token and self.eos_token_id is not None:
75 output.append(self.eos_token_id)
76 return output
77
78 def get_special_tokens_mask(
79 self,
80 token_ids_0: list[int],
81 token_ids_1: list[int] | None = None,
82 already_has_special_tokens: bool = False,
83 ) -> list[int]:
84 if already_has_special_tokens:
85 special_ids = set(self.all_special_ids)
86 return [1 if token_id in special_ids else 0 for token_id in token_ids_0]
87
88 mask = [0] * len(token_ids_0)
89 if self.add_bos_token and self.bos_token_id is not None:
90 mask = [1] + mask
91 if self.add_eos_token and self.eos_token_id is not None:
92 mask = mask + [1]
93 if token_ids_1 is not None:
94 mask += [0] * len(token_ids_1)
95 if self.add_eos_token and self.eos_token_id is not None:
96 mask.append(1)
97 return mask
98
99 def create_token_type_ids_from_sequences(
100 self,
101 token_ids_0: list[int],
102 token_ids_1: list[int] | None = None,
103 ) -> list[int]:
104 return [0] * len(self.build_inputs_with_special_tokens(token_ids_0, token_ids_1))
105
106 def save_vocabulary(self, save_directory: str, filename_prefix: str | None = None) -> tuple[str]:
107 output_name = "tokenizer.model"
108 if filename_prefix:
109 output_name = f"{filename_prefix}-{output_name}"
110 output_path = Path(save_directory) / output_name
111 if Path(self.vocab_file).resolve() != output_path.resolve():
112 shutil.copyfile(self.vocab_file, output_path)
113 return (str(output_path),)
114 