CoolFace
Modelpublic

timorobrecht/10k_initials_gpu1

sourceHugging Faceupdated 4mo agoView on Hugging Face
0likes2downloads
tokenization_pinyin_code.py114 linesDownload Raw Back to hf
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