AyaGL commited on
Commit
b6e0ef8
·
verified ·
1 Parent(s): 7e0a42d

Update tokenization_steerling.py

Browse files
Files changed (1) hide show
  1. tokenization_steerling.py +184 -0
tokenization_steerling.py ADDED
@@ -0,0 +1,184 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
+ from __future__ import annotations
2
+ from typing import Any
3
+ import tiktoken
4
+ from transformers import PreTrainedTokenizer
5
+
6
+ import tiktoken
7
+
8
+ class _SteerlingTokenizer:
9
+ """
10
+ Tokenizer for Steerling models.
11
+
12
+ Uses tiktoken cl100k_base with custom special tokens.
13
+ Pass ``instruct=True`` to include the 3 additional chat tokens
14
+ used by the instruct model.
15
+ """
16
+ ENCODING_NAME = 'cl100k_base'
17
+
18
+ def __init__(self, instruct: bool=False):
19
+ base_enc = tiktoken.get_encoding(self.ENCODING_NAME)
20
+ base_vocab = base_enc.n_vocab
21
+ self._pad_token_id = base_vocab
22
+ self._bos_token_id = base_vocab + 1
23
+ self._endofchunk_token_id = base_vocab + 2
24
+ self._mask_token_id = base_vocab + 3
25
+ self._eos_token_id = base_enc._special_tokens['<|endoftext|>']
26
+ self._instruct = instruct
27
+ special_tokens = {**base_enc._special_tokens, '<|pad|>': self._pad_token_id, '<|bos|>': self._bos_token_id, '<|endofchunk|>': self._endofchunk_token_id, '<|mask|>': self._mask_token_id}
28
+ if instruct:
29
+ self._start_header_id = base_vocab + 4
30
+ self._end_header_id = base_vocab + 5
31
+ self._eot_id = base_vocab + 6
32
+ self._vocab_size = base_vocab + 7
33
+ special_tokens.update({'<|start_header_id|>': self._start_header_id, '<|end_header_id|>': self._end_header_id, '<|eot_id|>': self._eot_id})
34
+ else:
35
+ self._start_header_id = None
36
+ self._end_header_id = None
37
+ self._eot_id = None
38
+ self._vocab_size = base_vocab + 4
39
+ self._tokenizer = tiktoken.Encoding(name=f'{self.ENCODING_NAME}_steerling', pat_str=base_enc._pat_str, mergeable_ranks=base_enc._mergeable_ranks, special_tokens=special_tokens)
40
+ self._special_token_ids = {self._pad_token_id, self._bos_token_id, self._eos_token_id, self._endofchunk_token_id, self._mask_token_id}
41
+ if instruct:
42
+ self._special_token_ids.update({self._start_header_id, self._end_header_id, self._eot_id})
43
+
44
+ def encode(self, text: str, add_special_tokens: bool=True) -> list[int]:
45
+ """
46
+ Encode text to token IDs.
47
+
48
+ Args:
49
+ text: Input text
50
+ add_special_tokens: If True, prepend BOS and append EOS
51
+
52
+ Returns:
53
+ List of token IDs
54
+ """
55
+ tokens = self._tokenizer.encode(text, disallowed_special=())
56
+ if add_special_tokens:
57
+ tokens = [self._bos_token_id] + tokens + [self._eos_token_id]
58
+ return tokens
59
+
60
+ def decode(self, tokens: list[int], skip_special_tokens: bool=True) -> str:
61
+ """
62
+ Decode token IDs to text.
63
+
64
+ Args:
65
+ tokens: Token IDs (list, numpy array, or torch tensor)
66
+ skip_special_tokens: If True, filter out special tokens before decoding
67
+
68
+ Returns:
69
+ Decoded text
70
+ """
71
+ if skip_special_tokens:
72
+ tokens = [int(t) for t in tokens if int(t) not in self._special_token_ids]
73
+ else:
74
+ tokens = [int(t) for t in tokens]
75
+ return self._tokenizer.decode(tokens)
76
+
77
+ @property
78
+ def vocab_size(self) -> int:
79
+ return self._vocab_size
80
+
81
+ @property
82
+ def pad_token_id(self) -> int:
83
+ return self._pad_token_id
84
+
85
+ @property
86
+ def bos_token_id(self) -> int:
87
+ return self._bos_token_id
88
+
89
+ @property
90
+ def eos_token_id(self) -> int:
91
+ return self._eos_token_id
92
+
93
+ @property
94
+ def endofchunk_token_id(self) -> int:
95
+ return self._endofchunk_token_id
96
+
97
+ @property
98
+ def mask_token_id(self) -> int:
99
+ return self._mask_token_id
100
+
101
+ @property
102
+ def instruct(self) -> bool:
103
+ return self._instruct
104
+
105
+ @property
106
+ def start_header_id(self) -> int | None:
107
+ return self._start_header_id
108
+
109
+ @property
110
+ def end_header_id(self) -> int | None:
111
+ return self._end_header_id
112
+
113
+ @property
114
+ def eot_id(self) -> int | None:
115
+ return self._eot_id
116
+
117
+ class SteerlingTokenizer(PreTrainedTokenizer):
118
+ vocab_files_names: dict[str, str] = {}
119
+ model_input_names = ["input_ids", "attention_mask"]
120
+
121
+ def __init__(self, encoding_name="cl100k_base", pad_token_id=100277,
122
+ bos_token_id=100278, eos_token_id=100257,
123
+ endofchunk_token_id=100279, mask_token_id=100280, **kwargs):
124
+ self._core = _SteerlingTokenizer(instruct=True)
125
+ self._endofchunk_token_id = endofchunk_token_id
126
+ self._mask_token_id = mask_token_id
127
+ for k in ("pad_token", "bos_token", "eos_token", "additional_special_tokens"):
128
+ kwargs.pop(k, None)
129
+ super().__init__(pad_token="<|pad|>", bos_token="<|bos|>", eos_token="<|endoftext|>",
130
+ additional_special_tokens=["<|endofchunk|>", "<|mask|>", "<|start_header_id|>", "<|end_header_id|>", "<|eot_id|>"], **kwargs)
131
+
132
+ @property
133
+ def vocab_size(self): return self._core.vocab_size
134
+ @property
135
+ def endofchunk_token_id(self): return self._core.endofchunk_token_id
136
+ @property
137
+ def mask_token_id(self): return self._core.mask_token_id
138
+
139
+ @property
140
+ def start_header_id(self): return self._core.start_header_id
141
+ @property
142
+ def end_header_id(self): return self._core.end_header_id
143
+ @property
144
+ def eot_id(self): return self._core.eot_id
145
+
146
+ def get_vocab(self): return dict(self._core._tokenizer._special_tokens)
147
+
148
+ def _tokenize(self, text, **kwargs):
149
+ return [str(i) for i in self._core._tokenizer.encode(text, disallowed_special=())]
150
+
151
+ def _convert_token_to_id(self, token):
152
+ special = self._core._tokenizer._special_tokens
153
+ if token in special: return special[token]
154
+ try: return int(token)
155
+ except ValueError:
156
+ ids = self._core._tokenizer.encode(token, disallowed_special=())
157
+ return ids[0] if ids else self._core.pad_token_id
158
+
159
+ def _convert_id_to_token(self, index):
160
+ for name, idx in self._core._tokenizer._special_tokens.items():
161
+ if idx == index: return name
162
+ try: return self._core._tokenizer.decode([index])
163
+ except Exception: return f"<|token_{index}|>"
164
+
165
+ def convert_tokens_to_string(self, tokens):
166
+ ids, special = [], self._core._tokenizer._special_tokens
167
+ for t in tokens:
168
+ if t in special: continue
169
+ try:
170
+ tid = int(t)
171
+ if tid not in self._core._special_token_ids: ids.append(tid)
172
+ except ValueError:
173
+ ids.extend(self._core._tokenizer.encode(t, disallowed_special=()))
174
+ return self._core._tokenizer.decode(ids)
175
+
176
+ def _decode(self, token_ids, skip_special_tokens=False, **kwargs):
177
+ return self._core.decode(list(token_ids) if not isinstance(token_ids, list) else token_ids,
178
+ skip_special_tokens=skip_special_tokens)
179
+
180
+ def build_inputs_with_special_tokens(self, token_ids_0, token_ids_1=None):
181
+ return token_ids_0
182
+
183
+ def save_vocabulary(self, save_directory, filename_prefix=None):
184
+ return ()