Arush kumar commited on
Commit
71f1a2a
·
1 Parent(s): cbc2499

Update tokenizer.py

Browse files
Files changed (1) hide show
  1. tokenizer.py +168 -282
tokenizer.py CHANGED
@@ -5,52 +5,62 @@ import logging
5
  from pathlib import Path
6
  from typing import Sequence, List, Optional
7
 
8
- from tokenizers import Tokenizer, pre_tokenizers, decoders, trainers
9
- from tokenizers.models import BPE
10
- from tokenizers.processors import TemplateProcessing
11
- import tiktoken
12
 
13
  logger = logging.getLogger(__name__)
14
 
15
- # GPT-2 style byte-level regex pretokenizer pattern. This is what actually
16
- # prevents "train." and "train!" from ever becoming distinct merged tokens
17
- # in the first place: punctuation, whitespace, and word characters are
18
- # split into separate pretoken chunks BEFORE BPE merges run, so BPE only
19
- # ever sees clean "train" as a candidate word-piece, with "." / "!" as
20
- # their own separate single-character pretokens. This is a pretokenization
21
- # property, not something achieved by pruning merges after training.
22
- GPT2_SPLIT_PATTERN = (
23
- r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""
24
- )
25
 
26
- SPECIAL_TOKENS = ["<unk>", "<pad>", "<bos>", "<eos>"]
27
- UNK_ID, PAD_ID, BOS_ID, EOS_ID = 0, 1, 2, 3
28
 
29
-
30
- def token_train(
31
  data_files: Sequence[str],
32
  model_prefix: str = 'tokenizer',
33
- vocab_size: int = 32000,
34
- min_frequency: int = 2,
 
 
 
 
 
 
 
 
 
 
 
 
 
 
35
  ) -> str:
36
  """
37
- Train a byte-level BPE tokenizer (tiktoken/GPT-2 style) using HF
38
- `tokenizers` as the trainer backend, then export a tiktoken-loadable
39
- merges file for fast inference.
40
-
41
- Why HF `tokenizers` as the trainer and not tiktoken directly: tiktoken
42
- itself is inference-only — it has no production-grade vocab trainer
43
- (its reference trainer in `_educational.py` is explicitly slow/
44
- reference-quality, not meant for real corpora). `tokenizers` is the
45
- standard fast Rust-backed BPE trainer; we train there, then hand off
46
- the resulting merges to tiktoken's `Encoding` for fast BPE inference
47
- at train/eval time.
48
-
49
- The GPT2_SPLIT_PATTERN pretokenizer is what keeps "train.", "train!",
50
- "train," etc. from ever being learned as single fused tokens — the
51
- regex splits words from adjacent punctuation before BPE ever runs, so
52
- BPE always sees "train" as a standalone candidate, with punctuation as
53
- its own separate byte-level pretoken.
 
 
 
 
 
 
 
54
  """
55
  data_files = [str(Path(p)) for p in data_files]
56
  if not data_files:
@@ -60,278 +70,153 @@ def token_train(
60
  if missing:
61
  raise FileNotFoundError(f'Missing input files: {missing}')
62
 
63
- tokenizer = Tokenizer(BPE(unk_token="<unk>", byte_fallback=True))
64
- tokenizer.pre_tokenizer = pre_tokenizers.ByteLevel(
65
- add_prefix_space=False, use_regex=True,
66
- )
67
- tokenizer.decoder = decoders.ByteLevel()
68
-
69
- trainer = trainers.BpeTrainer(
70
- vocab_size=vocab_size,
71
- min_frequency=min_frequency,
72
- special_tokens=SPECIAL_TOKENS,
73
- initial_alphabet=pre_tokenizers.ByteLevel.alphabet(),
74
- show_progress=True,
75
- )
76
-
77
- logger.info(f"Training BPE tokenizer: vocab_size={vocab_size} files={len(data_files)}")
78
- tokenizer.train(files=data_files, trainer=trainer)
79
-
80
- tokenizer.post_processor = TemplateProcessing(
81
- single="<bos> $A <eos>",
82
- special_tokens=[
83
- ("<bos>", tokenizer.token_to_id("<bos>")),
84
- ("<eos>", tokenizer.token_to_id("<eos>")),
85
- ],
86
  )
87
 
88
- hf_json_path = f'{model_prefix}.tokenizer.json'
89
- tokenizer.save(hf_json_path)
90
-
91
- tiktoken_path = _export_tiktoken_format(tokenizer, model_prefix)
92
- _validate_trained_model(hf_json_path, tiktoken_path, vocab_size)
93
-
94
- return tiktoken_path
95
-
96
-
97
- def _bytes_to_unicode() -> dict:
98
- """
99
- The canonical GPT-2 byte<->unicode bijection (same table used inside
100
- HF's ByteLevel pre_tokenizer/decoder, and in tiktoken's own reference
101
- implementations). Maps every raw byte value 0-255 to a printable
102
- unicode code point, so byte-level BPE can be trained/represented as
103
- ordinary text.
104
 
105
- Returns: dict {byte_value (0-255): unicode_char}
106
- """
107
- bs = (
108
- list(range(ord("!"), ord("~") + 1))
109
- + list(range(ord("¡"), ord("¬") + 1))
110
- + list(range(ord("®"), ord("ÿ") + 1))
111
  )
112
- cs = bs[:]
113
- n = 0
114
- for b in range(256):
115
- if b not in bs:
116
- bs.append(b)
117
- cs.append(256 + n)
118
- n += 1
119
- cs = [chr(c) for c in cs]
120
- return dict(zip(bs, cs))
121
-
122
-
123
- # Built once at import time: char -> original raw byte value. This is the
124
- # ONLY correct way to recover a token's raw bytes from its GPT-2
125
- # byte-level-alphabet string form.
126
- _UNICODE_TO_BYTE = {v: k for k, v in _bytes_to_unicode().items()}
127
-
128
-
129
- def _token_str_to_bytes(token_str: str) -> bytes:
130
- """
131
- Converts a GPT-2 byte-level-alphabet token string back to its raw
132
- bytes by inverting the bijection CHARACTER BY CHARACTER.
133
-
134
- CRITICAL: this must NOT go through `tokenizer.decoder.decode(...)`
135
- followed by `.encode('utf-8')`. That round-trip treats the decoded
136
- result as TEXT, and a lone raw byte in the 128-255 range is not valid
137
- UTF-8 on its own — decoding it in isolation silently produces the
138
- Unicode replacement character (U+FFFD) instead of the original byte.
139
- Every byte value 128-255 collapses into that same wrong 3-byte
140
- sequence this way, which is exactly what caused the "no entry found
141
- for key" Rust panic: 128 real single-byte entries were silently
142
- replaced by 127 colliding, wrong entries, deleting the entire
143
- non-ASCII byte range from the exported vocabulary. Any text
144
- containing so much as one accented character, curly quote, or emoji
145
- (i.e. almost any real-world corpus) then has no byte-fallback entry
146
- to encode with, and tiktoken's Rust core panics.
147
- """
148
- return bytes(_UNICODE_TO_BYTE[ch] for ch in token_str)
149
 
150
 
151
- def _export_tiktoken_format(tokenizer: "Tokenizer", model_prefix: str) -> str:
 
 
 
 
 
 
 
152
  """
153
- Export the trained HF tokenizer's merges into a tiktoken-loadable
154
- `.tiktoken` file (rank<TAB>base64-token-bytes per line) plus a sidecar
155
- JSON with special-token ids, so TokenizerWrapper can load via
156
- tiktoken.Encoding for fast inference.
157
-
158
- CRITICAL: tiktoken requires mergeable_ranks to be DENSE, starting at 0
159
- (rank IS the merge-priority order — see any real tiktoken.Encoding
160
- setup, e.g. Llama's tokenizer: special tokens are always assigned
161
- `len(mergeable_ranks) + i`, i.e. right after the merge ranks end).
162
- HF's BpeTrainer, because special_tokens is passed first, assigns them
163
- ids 0-3 and pushes actual merges to start at id 4 — a gapped,
164
- non-contiguous rank space if reused directly. Writing those original
165
- ids into the .tiktoken file silently corrupts BPE merge-priority
166
- ordering (tokenization no longer matches what was trained), and
167
- decode() maps ids to the wrong byte sequences — this was the actual
168
- cause of a real repetition/garbage-decode bug in production ("to to
169
- to", ",,,"), not a decoding-strategy issue. Fix: renumber merge ranks
170
- densely from 0 in sorted-original-id order (preserves relative merge
171
- priority, which is all that matters), then assign special token ids
172
- as len(mergeable_ranks) + i, matching tiktoken's actual convention.
173
  """
174
- import base64
175
-
176
- vocab = tokenizer.get_vocab() # token string -> original HF id
177
-
178
- # Sort NON-special tokens by their original HF id (preserves the
179
- # relative merge-priority order learned during training), then
180
- # renumber them densely 0..N-1 — this is the rank space tiktoken
181
- # actually requires.
182
- non_special = sorted(
183
- ((s, i) for s, i in vocab.items() if s not in SPECIAL_TOKENS),
184
- key=lambda kv: kv[1],
185
- )
186
-
187
- tiktoken_path = f'{model_prefix}.tiktoken'
188
- with open(tiktoken_path, 'w', encoding='utf-8') as f:
189
- for new_rank, (token_str, _old_id) in enumerate(non_special):
190
- # Direct char-by-char inversion of the GPT-2 byte<->unicode
191
- # bijection — NOT tokenizer.decoder.decode() + UTF-8 encode,
192
- # which is lossy for any byte value 128-255 (see
193
- # _token_str_to_bytes docstring for exactly why).
194
- token_bytes = _token_str_to_bytes(token_str)
195
- f.write(f"{base64.b64encode(token_bytes).decode('ascii')} {new_rank}\n")
196
-
197
- # Special tokens go right after the dense merge-rank space ends —
198
- # matching tiktoken's real convention (len(mergeable_ranks) + i), NOT
199
- # HF's original 0-3 ids.
200
- num_merges = len(non_special)
201
- special_ids = {
202
- name.strip('<>'): num_merges + i for i, name in enumerate(SPECIAL_TOKENS)
203
- }
204
- Path(f'{model_prefix}.special_tokens.json').write_text(
205
- json.dumps(special_ids, indent=2), encoding='utf-8'
206
- )
207
-
208
- # ── Hard completeness check ────────────────────────────────────────
209
- # tiktoken's BPE core requires every possible single byte (0-255) to
210
- # have a mergeable_ranks entry — it's the mandatory fallback base case
211
- # for encoding any byte the trained merges don't otherwise cover. If
212
- # even one byte value is missing (e.g. from a bug like the one this
213
- # function just fixed), encoding ANY text containing that byte panics
214
- # deep inside tiktoken's Rust core with a cryptic "no entry found for
215
- # key" — with no indication of which byte or why. Catch that here,
216
- # immediately after export, with a clear Python-level error instead.
217
- exported_single_bytes = set()
218
- with open(tiktoken_path, 'r', encoding='utf-8') as f:
219
- for line in f:
220
- line = line.strip()
221
- if not line:
222
- continue
223
- b64_token, _ = line.split()
224
- raw = base64.b64decode(b64_token)
225
- if len(raw) == 1:
226
- exported_single_bytes.add(raw[0])
227
- missing = sorted(set(range(256)) - exported_single_bytes)
228
- if missing:
229
- raise RuntimeError(
230
- f"Tokenizer export is missing {len(missing)}/256 single-byte "
231
- f"entries (byte values: {missing[:20]}{'...' if len(missing) > 20 else ''}). "
232
- f"Any training text containing these byte values will crash "
233
- f"tiktoken's Rust core with 'no entry found for key'. This "
234
- f"means the byte<->unicode inversion during export is broken — "
235
- f"do not proceed with this .tiktoken file."
236
- )
237
-
238
- return tiktoken_path
239
-
240
 
241
- def _validate_trained_model(hf_json_path: str, tiktoken_path: str, expected_vocab_size: int) -> None:
242
- """Round-trip check on the HF tokenizer (source of truth) before we
243
- trust the exported tiktoken artifact."""
244
- tok = Tokenizer.from_file(hf_json_path)
245
- actual_vocab = tok.get_vocab_size()
246
  if actual_vocab != expected_vocab_size:
247
- logger.warning(f"Trained vocab_size={actual_vocab} differs from requested={expected_vocab_size}")
248
-
249
- for name in SPECIAL_TOKENS:
250
- if tok.token_to_id(name) is None:
251
- raise ValueError(f'Trained model missing special token {name}')
252
-
253
- probe = "The quick brown fox jumps over 42 lazy dogs. Did it train? train! train."
254
- encoding = tok.encode(probe)
255
- if not encoding.ids:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
256
  raise ValueError('Validation encode produced empty output')
257
- decoded = tok.decode(encoding.ids)
258
  if not decoded.strip():
259
  raise ValueError('Validation round-trip produced empty decode')
260
 
261
- logger.info(f"Validation OK: vocab={actual_vocab} probe_tokens={len(encoding.ids)}")
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
262
 
 
 
263
 
264
- class TokenizerWrapper:
265
- """
266
- Fast-path wrapper using tiktoken for encode/decode at train/eval time,
267
- backed by a vocab trained via HF `tokenizers` (see train_bpe_tokenizer).
268
- """
269
 
 
270
  def __init__(self, model_path: str):
271
  model_path = str(Path(model_path))
272
- tiktoken_path = model_path if model_path.endswith('.tiktoken') else f'{model_path}.tiktoken'
273
- special_path = tiktoken_path.replace('.tiktoken', '.special_tokens.json')
274
-
275
- if not Path(tiktoken_path).exists():
276
- raise FileNotFoundError(tiktoken_path)
277
- if not Path(special_path).exists():
278
- raise FileNotFoundError(special_path)
279
-
280
- special_ids = json.loads(Path(special_path).read_text(encoding='utf-8'))
281
- self.unk_id = special_ids['unk']
282
- self.pad_id = special_ids['pad']
283
- self.bos_id = special_ids['bos']
284
- self.eos_id = special_ids['eos']
285
  for name, val in [('pad', self.pad_id), ('unk', self.unk_id), ('bos', self.bos_id), ('eos', self.eos_id)]:
286
- if val is None or val < 0:
287
- raise ValueError(f'Tokenizer missing <{name}>')
288
-
289
- mergeable_ranks = self._load_tiktoken_ranks(tiktoken_path)
290
- self.vocab_size = len(mergeable_ranks) + len(special_ids)
291
-
292
- self.enc = tiktoken.Encoding(
293
- name=Path(tiktoken_path).stem,
294
- pat_str=GPT2_SPLIT_PATTERN,
295
- mergeable_ranks=mergeable_ranks,
296
- special_tokens={
297
- "<unk>": self.unk_id, "<pad>": self.pad_id,
298
- "<bos>": self.bos_id, "<eos>": self.eos_id,
299
- },
300
- )
301
  self._special_ids = {self.pad_id, self.bos_id, self.eos_id}
302
 
303
- @staticmethod
304
- def _load_tiktoken_ranks(tiktoken_path: str) -> dict:
305
- ranks = {}
306
- with open(tiktoken_path, 'r', encoding='utf-8') as f:
307
- for line in f:
308
- line = line.strip()
309
- if not line:
310
- continue
311
- b64_token, rank = line.split()
312
- import base64
313
- ranks[base64.b64decode(b64_token)] = int(rank)
314
- return ranks
315
-
316
  def encode(self, text: str, add_bos: bool = True, add_eos: bool = False) -> List[int]:
317
  if text is None:
318
  raise ValueError('encode() received None')
319
  if text == '':
320
  ids: List[int] = []
321
  else:
322
- # disallowed_special=() (not allowed_special=set()): training
323
- # corpora can contain LITERAL text that happens to match a
324
- # special token string — e.g. WikiText-103 ships with literal
325
- # "<unk>" markers baked into its raw text from its original
326
- # preprocessing. tiktoken's default behavior treats any string
327
- # matching a registered special token as a forbidden injection
328
- # attempt and raises. We want the opposite: literal "<unk>" in
329
- # training data should be encoded as ordinary text (broken into
330
- # bytes/subwords), never treated as an actual special token
331
- # unless WE insert it programmatically (e.g. via add_bos/add_eos
332
- # below, which append the token ID directly, bypassing string
333
- # matching entirely — so this doesn't weaken bos/eos handling).
334
- ids = self.enc.encode(text, disallowed_special=())
335
  if add_bos:
336
  ids = [self.bos_id] + ids
337
  if add_eos:
@@ -361,7 +246,7 @@ class TokenizerWrapper:
361
  filtered = [int(i) for i in ids if int(i) not in self._special_ids]
362
  else:
363
  filtered = [int(i) for i in ids if int(i) != self.pad_id]
364
- return self.enc.decode(filtered)
365
 
366
  def decode_batch(self, batch_ids: Sequence[Sequence[int]], skip_special_tokens: bool = True) -> List[str]:
367
  return [self.decode(ids, skip_special_tokens=skip_special_tokens) for ids in batch_ids]
@@ -377,6 +262,7 @@ class TokenizerWrapper:
377
 
378
  @classmethod
379
  def from_config(cls, model_path: str, config_path: Optional[str] = None) -> 'TokenizerWrapper':
 
380
  tok = cls(model_path)
381
  if config_path and Path(config_path).exists():
382
  cfg = json.loads(Path(config_path).read_text(encoding='utf-8'))
 
5
  from pathlib import Path
6
  from typing import Sequence, List, Optional
7
 
8
+ import sentencepiece as spm
 
 
 
9
 
10
  logger = logging.getLogger(__name__)
11
 
12
+ # 32K is the well-established baseline vocab size for BPE/SentencePiece
13
+ # LLM tokenizers (Llama-1/2, T5, Gopher, Chinchilla all use exactly this).
14
+ # 128K+ only pays off for heavy multilingual/code coverage; for a small,
15
+ # largely-English, narrow-domain model, 32K is the standard, safe default.
16
+ DEFAULT_VOCAB_SIZE = 32000
 
 
 
 
 
17
 
 
 
18
 
19
+ def train_sentencepiece(
 
20
  data_files: Sequence[str],
21
  model_prefix: str = 'tokenizer',
22
+ vocab_size: int = DEFAULT_VOCAB_SIZE,
23
+ model_type: str = 'bpe',
24
+ character_coverage: float = 0.9995,
25
+ byte_fallback: bool = True,
26
+ pad_id: int = 1,
27
+ unk_id: int = 0,
28
+ bos_id: int = 2,
29
+ eos_id: int = 3,
30
+ add_dummy_prefix: bool = True,
31
+ num_threads: int = 8,
32
+ input_sentence_size: int = 5_000_000,
33
+ shuffle_input_sentence: bool = True,
34
+ max_sentence_length: int = 16384,
35
+ split_digits: bool = True,
36
+ allow_whitespace_only_pieces: bool = True,
37
+ train_extremely_large_corpus: bool = False,
38
  ) -> str:
39
  """
40
+ Train a SentencePiece BPE tokenizer with byte-fallback — the same
41
+ scheme used by Llama-2, Mistral, and EuroLLM (BPE + byte_fallback via
42
+ SentencePiece specifically, not a hand-rolled BPE implementation).
43
+
44
+ Why SentencePiece and not a hand-written tiktoken export: SentencePiece's
45
+ C++ core does encode/decode and merge-rank bookkeeping internally and
46
+ natively — there is no manual ID-renumbering or rank-export step for
47
+ calling code to get wrong. (A prior tiktoken-based rewrite of this
48
+ tokenizer had exactly that class of bug: hand-exported merge ranks were
49
+ non-contiguous because special tokens occupied ids 0-3 in the source
50
+ vocab, silently corrupting merge-priority order and decode() mappings —
51
+ manifesting as repetitive garbage output like "to to to" despite a
52
+ healthy training loss. Delegating to SentencePiece's own encode/decode
53
+ removes that entire class of bug by construction.)
54
+
55
+ Notes on defaults:
56
+ - character_coverage < 1.0 with byte_fallback=True: rare glyphs fall
57
+ back to byte pieces instead of bloating the vocab with singletons.
58
+ - input_sentence_size + shuffle_input_sentence: without shuffling,
59
+ SentencePiece samples from the START of the concatenated corpus,
60
+ which silently biases vocab toward whichever domain file comes
61
+ first if you hand it multiple files back to back.
62
+ - split_digits: keeps numbers as individual digit tokens, which
63
+ generally helps arithmetic/math task tokenization consistency.
64
  """
65
  data_files = [str(Path(p)) for p in data_files]
66
  if not data_files:
 
70
  if missing:
71
  raise FileNotFoundError(f'Missing input files: {missing}')
72
 
73
+ kwargs = dict(
74
+ input=','.join(data_files),
75
+ model_prefix=model_prefix,
76
+ vocab_size=int(vocab_size),
77
+ model_type=model_type,
78
+ character_coverage=character_coverage,
79
+ pad_id=pad_id,
80
+ unk_id=unk_id,
81
+ bos_id=bos_id,
82
+ eos_id=eos_id,
83
+ byte_fallback=byte_fallback,
84
+ hard_vocab_limit=False,
85
+ normalization_rule_name='nmt_nfkc',
86
+ add_dummy_prefix=add_dummy_prefix,
87
+ num_threads=num_threads,
88
+ input_sentence_size=input_sentence_size,
89
+ shuffle_input_sentence=shuffle_input_sentence,
90
+ max_sentence_length=max_sentence_length,
91
+ split_digits=split_digits,
92
+ allow_whitespace_only_pieces=allow_whitespace_only_pieces,
93
+ train_extremely_large_corpus=train_extremely_large_corpus,
 
 
94
  )
95
 
96
+ logger.info(f"Training SentencePiece: vocab_size={vocab_size} model_type={model_type} "
97
+ f"files={len(data_files)}")
98
+ spm.SentencePieceTrainer.train(**kwargs)
 
 
 
 
 
 
 
 
 
 
 
 
 
99
 
100
+ model_path = f'{model_prefix}.model'
101
+ _validate_trained_model(
102
+ model_path, vocab_size,
103
+ expected_pad=pad_id, expected_unk=unk_id, expected_bos=bos_id, expected_eos=eos_id,
 
 
104
  )
105
+ return model_path
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
106
 
107
 
108
+ def _validate_trained_model(
109
+ model_path: str,
110
+ expected_vocab_size: int,
111
+ expected_pad: int,
112
+ expected_unk: int,
113
+ expected_bos: int,
114
+ expected_eos: int,
115
+ ) -> None:
116
  """
117
+ Self-critique validation pass — checks the things that actually broke
118
+ in the previous (tiktoken) tokenizer, not just "does it load".
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
119
  """
120
+ sp = spm.SentencePieceProcessor(model_file=model_path)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
121
 
122
+ # 1. Vocab size sanity
123
+ actual_vocab = sp.vocab_size()
 
 
 
124
  if actual_vocab != expected_vocab_size:
125
+ logger.warning(f"Trained vocab_size={actual_vocab} differs from requested={expected_vocab_size} "
126
+ f"(hard_vocab_limit=False allows this if the corpus is small)")
127
+
128
+ # 2. Special token IDs must be EXACTLY what was requested — not just
129
+ # ">= 0". A previous bug class involved special-token ids silently
130
+ # drifting from what calling code assumed. Check explicitly, not
131
+ # loosely.
132
+ checks = [
133
+ ('pad', sp.pad_id(), expected_pad),
134
+ ('unk', sp.unk_id(), expected_unk),
135
+ ('bos', sp.bos_id(), expected_bos),
136
+ ('eos', sp.eos_id(), expected_eos),
137
+ ]
138
+ for name, actual, expected in checks:
139
+ if actual < 0:
140
+ raise ValueError(f'Trained model missing <{name}> special token')
141
+ if actual != expected:
142
+ raise ValueError(
143
+ f'<{name}> id drift: requested {expected}, SentencePiece '
144
+ f'assigned {actual}. This mismatch is exactly the class of '
145
+ f'bug that broke a previous tokenizer version — refusing '
146
+ f'to silently proceed.'
147
+ )
148
+
149
+ # 3. Basic round-trip: encode -> decode must reproduce recognizable text
150
+ probe = "The quick brown fox jumps over 42 lazy dogs. def foo(): return None"
151
+ ids = sp.encode(probe, out_type=int)
152
+ if not ids:
153
  raise ValueError('Validation encode produced empty output')
154
+ decoded = sp.decode(ids)
155
  if not decoded.strip():
156
  raise ValueError('Validation round-trip produced empty decode')
157
 
158
+ # 4. SPECIFIC regression check for the actual reported failure mode:
159
+ # repetitive-token degenerate decode ("to to to", ",,,"). This won't
160
+ # catch a MODEL that's actually stuck in a repetition loop (that's a
161
+ # decoding-strategy issue, separate from the tokenizer), but it DOES
162
+ # catch a tokenizer that maps distinct ids to the same or corrupted
163
+ # text, which was the real bug here: encode the same repeated-word
164
+ # probe multiple times and confirm token ids are stable and decode
165
+ # is exact, not degenerating into duplicated/garbled pieces.
166
+ repeat_probe = "to to to , , , the the the"
167
+ repeat_ids = sp.encode(repeat_probe, out_type=int)
168
+ repeat_decoded = sp.decode(repeat_ids)
169
+ # Re-encoding the decoded output should reproduce the same ids
170
+ # (idempotency) — this is the real symptom check: a corrupted rank/id
171
+ # mapping breaks exactly this property even when a single encode/decode
172
+ # pass looks fine.
173
+ reencoded_ids = sp.encode(repeat_decoded, out_type=int)
174
+ if reencoded_ids != repeat_ids:
175
+ raise ValueError(
176
+ f'Round-trip idempotency FAILED on repeated-token probe: '
177
+ f'encode->decode->encode did not reproduce the same ids. '
178
+ f'original={repeat_ids} reencoded={reencoded_ids}. This is '
179
+ f'the specific failure signature of an id/rank mapping bug.'
180
+ )
181
+
182
+ # 5. Byte-fallback sanity: an unusual/rare unicode character must not
183
+ # crash and must not silently become <unk> if byte_fallback is on —
184
+ # it should decompose into byte pieces instead.
185
+ exotic_probe = "emoji test \U0001F600 and rare char \u0800"
186
+ exotic_ids = sp.encode(exotic_probe, out_type=int)
187
+ if not exotic_ids:
188
+ raise ValueError('Byte-fallback validation: exotic-character probe produced empty encode')
189
+ exotic_decoded = sp.decode(exotic_ids)
190
+ if not exotic_decoded.strip():
191
+ raise ValueError('Byte-fallback validation: exotic-character round-trip produced empty decode')
192
 
193
+ logger.info(f"✓ Validation OK: vocab={actual_vocab} probe_tokens={len(ids)} "
194
+ f"round-trip idempotency verified, byte-fallback verified")
195
 
 
 
 
 
 
196
 
197
+ class TokenizerWrapper:
198
  def __init__(self, model_path: str):
199
  model_path = str(Path(model_path))
200
+ if not Path(model_path).exists():
201
+ raise FileNotFoundError(model_path)
202
+ self.sp = spm.SentencePieceProcessor(model_file=model_path)
203
+ self.vocab_size = int(self.sp.vocab_size())
204
+ self.pad_id = self.sp.pad_id()
205
+ self.unk_id = self.sp.unk_id()
206
+ self.bos_id = self.sp.bos_id()
207
+ self.eos_id = self.sp.eos_id()
 
 
 
 
 
208
  for name, val in [('pad', self.pad_id), ('unk', self.unk_id), ('bos', self.bos_id), ('eos', self.eos_id)]:
209
+ if val < 0:
210
+ raise ValueError(f'SentencePiece model missing <{name}>')
 
 
 
 
 
 
 
 
 
 
 
 
 
211
  self._special_ids = {self.pad_id, self.bos_id, self.eos_id}
212
 
 
 
 
 
 
 
 
 
 
 
 
 
 
213
  def encode(self, text: str, add_bos: bool = True, add_eos: bool = False) -> List[int]:
214
  if text is None:
215
  raise ValueError('encode() received None')
216
  if text == '':
217
  ids: List[int] = []
218
  else:
219
+ ids = list(self.sp.encode(text, out_type=int))
 
 
 
 
 
 
 
 
 
 
 
 
220
  if add_bos:
221
  ids = [self.bos_id] + ids
222
  if add_eos:
 
246
  filtered = [int(i) for i in ids if int(i) not in self._special_ids]
247
  else:
248
  filtered = [int(i) for i in ids if int(i) != self.pad_id]
249
+ return self.sp.decode(filtered)
250
 
251
  def decode_batch(self, batch_ids: Sequence[Sequence[int]], skip_special_tokens: bool = True) -> List[str]:
252
  return [self.decode(ids, skip_special_tokens=skip_special_tokens) for ids in batch_ids]
 
262
 
263
  @classmethod
264
  def from_config(cls, model_path: str, config_path: Optional[str] = None) -> 'TokenizerWrapper':
265
+ """Load and, if a config is given, verify special-id consistency against it."""
266
  tok = cls(model_path)
267
  if config_path and Path(config_path).exists():
268
  cfg = json.loads(Path(config_path).read_text(encoding='utf-8'))