Implemented robust type handling and tokenizer exception safety

#2
Files changed (1) hide show
  1. tokenization_kimi.py +42 -20
tokenization_kimi.py CHANGED
@@ -94,9 +94,16 @@ class TikTokenTokenizer(PreTrainedTokenizer):
94
  "<|im_middle|>",
95
  ]
96
 
97
- special_tokens_mapping = {
98
- i: added_tokens_decoder[i].content for i in added_tokens_decoder
99
- }
 
 
 
 
 
 
 
100
 
101
  self.vocab_file = vocab_file
102
  mergeable_ranks = load_tiktoken_bpe(vocab_file)
@@ -108,8 +115,6 @@ class TikTokenTokenizer(PreTrainedTokenizer):
108
  )
109
  }
110
 
111
-
112
-
113
  self.model = tiktoken.Encoding(
114
  name=Path(vocab_file).name,
115
  pat_str=self.pat_str,
@@ -119,27 +124,39 @@ class TikTokenTokenizer(PreTrainedTokenizer):
119
  logger.info(f"Reloaded tiktoken model from {vocab_file}")
120
 
121
  self.n_words: int = self.model.n_vocab
122
- # BOS / EOS token IDs
123
- self.bos_id: int = self.special_tokens[str(bos_token)]
124
- self.eos_id: int = self.special_tokens[str(eos_token)]
 
 
 
 
 
 
 
 
 
125
  logger.info(
126
  f"#words: {self.n_words} - BOS ID: {self.bos_id} - EOS ID: {self.eos_id}"
127
  )
128
 
129
- self.pad_id: int = self.special_tokens[str(pad_token)]
130
- self.unk_id: int = self.special_tokens[str(unk_token)]
131
-
132
  self.byte_encoder = bytes_to_unicode()
133
  self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
134
 
135
  self.decoder = {}
 
136
  for i in range(self.n_words):
137
- # Taken from https://gist.github.com/xenova/a452a6474428de0182b17605a98631ee
138
- decoding = ''.join([
139
- self.byte_encoder[ord(char)] for char in
140
- self.model.decode_single_token_bytes(i).decode('latin-1')
141
- ])
142
- self.decoder[i] = decoding
 
 
 
 
 
143
 
144
  self.encoder = {}
145
  for i in range(self.n_words):
@@ -180,7 +197,7 @@ class TikTokenTokenizer(PreTrainedTokenizer):
180
  logger.warning( f"Calling super().encode with {kwargs}" )
181
  return super().encode(text, **kwargs)
182
 
183
- assert type(text) is str
184
 
185
  # The tiktoken tokenizer can handle <=400k chars without
186
  # pyo3_runtime.PanicException.
@@ -244,8 +261,13 @@ class TikTokenTokenizer(PreTrainedTokenizer):
244
  if len(kwargs) > 0:
245
  return super().decode(token_ids, **kwargs)
246
 
247
- if type(token_ids) is int:
248
- token_ids = [token_ids]
 
 
 
 
 
249
 
250
  return self.model.decode(cast(List[int], token_ids))
251
 
 
94
  "<|im_middle|>",
95
  ]
96
 
97
+ special_tokens_mapping = {}
98
+ if added_tokens_decoder is not None:
99
+ for k, v in added_tokens_decoder.items():
100
+ if isinstance(v, dict):
101
+ content = v.get("content", "")
102
+ elif hasattr(v, "content"):
103
+ content = v.content
104
+ else:
105
+ content = str(v)
106
+ special_tokens_mapping[int(k)] = content
107
 
108
  self.vocab_file = vocab_file
109
  mergeable_ranks = load_tiktoken_bpe(vocab_file)
 
115
  )
116
  }
117
 
 
 
118
  self.model = tiktoken.Encoding(
119
  name=Path(vocab_file).name,
120
  pat_str=self.pat_str,
 
124
  logger.info(f"Reloaded tiktoken model from {vocab_file}")
125
 
126
  self.n_words: int = self.model.n_vocab
127
+
128
+ # BOS / EOS / PAD / UNK string representation mapping
129
+ bos_str = bos_token.content if hasattr(bos_token, "content") else str(bos_token) if bos_token is not None else None
130
+ eos_str = eos_token.content if hasattr(eos_token, "content") else str(eos_token) if eos_token is not None else None
131
+ pad_str = pad_token.content if hasattr(pad_token, "content") else str(pad_token) if pad_token is not None else None
132
+ unk_str = unk_token.content if hasattr(unk_token, "content") else str(unk_token) if unk_token is not None else None
133
+
134
+ self.bos_id: Optional[int] = self.special_tokens.get(bos_str) if bos_str is not None else None
135
+ self.eos_id: Optional[int] = self.special_tokens.get(eos_str) if eos_str is not None else None
136
+ self.pad_id: Optional[int] = self.special_tokens.get(pad_str) if pad_str is not None else None
137
+ self.unk_id: Optional[int] = self.special_tokens.get(unk_str) if unk_str is not None else None
138
+
139
  logger.info(
140
  f"#words: {self.n_words} - BOS ID: {self.bos_id} - EOS ID: {self.eos_id}"
141
  )
142
 
 
 
 
143
  self.byte_encoder = bytes_to_unicode()
144
  self.byte_decoder = {v: k for k, v in self.byte_encoder.items()}
145
 
146
  self.decoder = {}
147
+ id_to_special = {v: k for k, v in self.special_tokens.items()}
148
  for i in range(self.n_words):
149
+ if i in id_to_special:
150
+ self.decoder[i] = id_to_special[i]
151
+ else:
152
+ try:
153
+ decoding = ''.join([
154
+ self.byte_encoder[ord(char)] for char in
155
+ self.model.decode_single_token_bytes(i).decode('latin-1')
156
+ ])
157
+ self.decoder[i] = decoding
158
+ except Exception:
159
+ self.decoder[i] = f"<|reserved_token_{i}|>"
160
 
161
  self.encoder = {}
162
  for i in range(self.n_words):
 
197
  logger.warning( f"Calling super().encode with {kwargs}" )
198
  return super().encode(text, **kwargs)
199
 
200
+ assert isinstance(text, str), f"text must be a string, got {type(text)}"
201
 
202
  # The tiktoken tokenizer can handle <=400k chars without
203
  # pyo3_runtime.PanicException.
 
261
  if len(kwargs) > 0:
262
  return super().decode(token_ids, **kwargs)
263
 
264
+ if hasattr(token_ids, "tolist"):
265
+ token_ids = token_ids.tolist()
266
+
267
+ if isinstance(token_ids, (int, float)):
268
+ token_ids = [int(token_ids)]
269
+ else:
270
+ token_ids = [int(x) for x in token_ids]
271
 
272
  return self.model.decode(cast(List[int], token_ids))
273