ArushBuilds commited on
Commit
ebd36d6
·
1 Parent(s): a3032cf

Update tokenizer.py

Browse files
Files changed (1) hide show
  1. tokenizer.py +311 -298
tokenizer.py CHANGED
@@ -1,283 +1,175 @@
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
1
  from __future__ import annotations
2
 
3
  import json
4
  import logging
 
 
5
  from pathlib import Path
6
- from typing import Sequence, List, Optional
7
 
8
  import numpy as np
9
- import sentencepiece as spm
10
-
11
- DEFAULT_CHUNK_CHARS = 2_000_000
12
 
13
  logger = logging.getLogger(__name__)
14
 
15
- # 32K is the well-established baseline vocab size for BPE/SentencePiece
16
- # LLM tokenizers (Llama-1/2, T5, Gopher, Chinchilla all use exactly this).
17
- # 128K+ only pays off for heavy multilingual/code coverage; for a small,
18
- # largely-English, narrow-domain model, 32K is the standard, safe default.
19
- DEFAULT_VOCAB_SIZE = 32000
20
-
21
-
22
- def train_sentencepiece(
23
- data_files: Sequence[str],
24
- model_prefix: str = 'tokenizer',
25
- vocab_size: int = DEFAULT_VOCAB_SIZE,
26
- model_type: str = 'bpe',
27
- character_coverage: float = 0.9995,
28
- byte_fallback: bool = True,
29
- pad_id: int = 1,
30
- unk_id: int = 0,
31
- bos_id: int = 2,
32
- eos_id: int = 3,
33
- add_dummy_prefix: bool = True,
34
- num_threads: int = 8,
35
- input_sentence_size: int = 5_000_000,
36
- shuffle_input_sentence: bool = True,
37
- max_sentence_length: int = 16384,
38
- split_digits: bool = True,
39
- allow_whitespace_only_pieces: bool = True,
40
- train_extremely_large_corpus: bool = False,
41
- ) -> str:
42
- """
43
- Train a SentencePiece BPE tokenizer with byte-fallback — the same
44
- scheme used by Llama-2, Mistral, and EuroLLM (BPE + byte_fallback via
45
- SentencePiece specifically, not a hand-rolled BPE implementation).
46
-
47
- Why SentencePiece and not a hand-written tiktoken export: SentencePiece's
48
- C++ core does encode/decode and merge-rank bookkeeping internally and
49
- natively — there is no manual ID-renumbering or rank-export step for
50
- calling code to get wrong. (A prior tiktoken-based rewrite of this
51
- tokenizer had exactly that class of bug: hand-exported merge ranks were
52
- non-contiguous because special tokens occupied ids 0-3 in the source
53
- vocab, silently corrupting merge-priority order and decode() mappings —
54
- manifesting as repetitive garbage output like "to to to" despite a
55
- healthy training loss. Delegating to SentencePiece's own encode/decode
56
- removes that entire class of bug by construction.)
57
-
58
- Notes on defaults:
59
- - character_coverage < 1.0 with byte_fallback=True: rare glyphs fall
60
- back to byte pieces instead of bloating the vocab with singletons.
61
- - input_sentence_size + shuffle_input_sentence: without shuffling,
62
- SentencePiece samples from the START of the concatenated corpus,
63
- which silently biases vocab toward whichever domain file comes
64
- first if you hand it multiple files back to back.
65
- - split_digits: keeps numbers as individual digit tokens, which
66
- generally helps arithmetic/math task tokenization consistency.
67
- """
68
- data_files = [str(Path(p)) for p in data_files]
69
- if not data_files:
70
- raise ValueError('data_files is empty')
71
-
72
- missing = [f for f in data_files if not Path(f).exists()]
73
- if missing:
74
- raise FileNotFoundError(f'Missing input files: {missing}')
75
-
76
- kwargs = dict(
77
- input=','.join(data_files),
78
- model_prefix=model_prefix,
79
- vocab_size=int(vocab_size),
80
- model_type=model_type,
81
- character_coverage=character_coverage,
82
- pad_id=pad_id,
83
- unk_id=unk_id,
84
- bos_id=bos_id,
85
- eos_id=eos_id,
86
- byte_fallback=byte_fallback,
87
- hard_vocab_limit=False,
88
- normalization_rule_name='nmt_nfkc',
89
- add_dummy_prefix=add_dummy_prefix,
90
- num_threads=num_threads,
91
- input_sentence_size=input_sentence_size,
92
- shuffle_input_sentence=shuffle_input_sentence,
93
- max_sentence_length=max_sentence_length,
94
- split_digits=split_digits,
95
- allow_whitespace_only_pieces=allow_whitespace_only_pieces,
96
- train_extremely_large_corpus=train_extremely_large_corpus,
97
- )
98
-
99
- logger.info(f"Training SentencePiece: vocab_size={vocab_size} model_type={model_type} "
100
- f"files={len(data_files)}")
101
- spm.SentencePieceTrainer.train(**kwargs)
102
-
103
- model_path = f'{model_prefix}.model'
104
- _validate_trained_model(
105
- model_path, vocab_size,
106
- expected_pad=pad_id, expected_unk=unk_id, expected_bos=bos_id, expected_eos=eos_id,
107
- )
108
- return model_path
109
-
110
-
111
- def _validate_trained_model(
112
- model_path: str,
113
- expected_vocab_size: int,
114
- expected_pad: int,
115
- expected_unk: int,
116
- expected_bos: int,
117
- expected_eos: int,
118
- ) -> None:
119
- """
120
- Self-critique validation pass — checks the things that actually broke
121
- in the previous (tiktoken) tokenizer, not just "does it load".
122
- """
123
- sp = spm.SentencePieceProcessor(model_file=model_path)
124
-
125
- # 1. Vocab size sanity
126
- actual_vocab = sp.vocab_size()
127
- if actual_vocab != expected_vocab_size:
128
- logger.warning(f"Trained vocab_size={actual_vocab} differs from requested={expected_vocab_size} "
129
- f"(hard_vocab_limit=False allows this if the corpus is small)")
130
-
131
- # 2. Special token IDs must be EXACTLY what was requested — not just
132
- # ">= 0". A previous bug class involved special-token ids silently
133
- # drifting from what calling code assumed. Check explicitly, not
134
- # loosely.
135
- checks = [
136
- ('pad', sp.pad_id(), expected_pad),
137
- ('unk', sp.unk_id(), expected_unk),
138
- ('bos', sp.bos_id(), expected_bos),
139
- ('eos', sp.eos_id(), expected_eos),
140
- ]
141
- for name, actual, expected in checks:
142
- if actual < 0:
143
- raise ValueError(f'Trained model missing <{name}> special token')
144
- if actual != expected:
145
- raise ValueError(
146
- f'<{name}> id drift: requested {expected}, SentencePiece '
147
- f'assigned {actual}. This mismatch is exactly the class of '
148
- f'bug that broke a previous tokenizer version — refusing '
149
- f'to silently proceed.'
150
  )
 
 
 
 
151
 
152
- # 3. Basic round-trip: encode -> decode must reproduce recognizable text
153
- probe = "The quick brown fox jumps over 42 lazy dogs. def foo(): return None"
154
- ids = sp.encode(probe, out_type=int)
155
- if not ids:
156
- raise ValueError('Validation encode produced empty output')
157
- decoded = sp.decode(ids)
158
- if not decoded.strip():
159
- raise ValueError('Validation round-trip produced empty decode')
160
-
161
- # 4. SPECIFIC regression check for the actual reported failure mode:
162
- # repetitive-token degenerate decode ("to to to", ",,,"). This won't
163
- # catch a MODEL that's actually stuck in a repetition loop (that's a
164
- # decoding-strategy issue, separate from the tokenizer), but it DOES
165
- # catch a tokenizer that maps distinct ids to the same or corrupted
166
- # text, which was the real bug here: encode the same repeated-word
167
- # probe multiple times and confirm token ids are stable and decode
168
- # is exact, not degenerating into duplicated/garbled pieces.
169
- repeat_probe = "to to to , , , the the the"
170
- repeat_ids = sp.encode(repeat_probe, out_type=int)
171
- repeat_decoded = sp.decode(repeat_ids)
172
- # Re-encoding the decoded output should reproduce the same ids
173
- # (idempotency) — this is the real symptom check: a corrupted rank/id
174
- # mapping breaks exactly this property even when a single encode/decode
175
- # pass looks fine.
176
- reencoded_ids = sp.encode(repeat_decoded, out_type=int)
177
- if reencoded_ids != repeat_ids:
178
- raise ValueError(
179
- f'Round-trip idempotency FAILED on repeated-token probe: '
180
- f'encode->decode->encode did not reproduce the same ids. '
181
- f'original={repeat_ids} reencoded={reencoded_ids}. This is '
182
- f'the specific failure signature of an id/rank mapping bug.'
183
- )
184
-
185
- # 5. Byte-fallback sanity: an unusual/rare unicode character must not
186
- # crash and must not silently become <unk> if byte_fallback is on —
187
- # it should decompose into byte pieces instead.
188
- exotic_probe = "emoji test \U0001F600 and rare char \u0800"
189
- exotic_ids = sp.encode(exotic_probe, out_type=int)
190
- if not exotic_ids:
191
- raise ValueError('Byte-fallback validation: exotic-character probe produced empty encode')
192
- exotic_decoded = sp.decode(exotic_ids)
193
- if not exotic_decoded.strip():
194
- raise ValueError('Byte-fallback validation: exotic-character round-trip produced empty decode')
195
-
196
- logger.info(f"✓ Validation OK: vocab={actual_vocab} probe_tokens={len(ids)} "
197
- f"round-trip idempotency verified, byte-fallback verified")
198
 
 
 
 
199
 
200
  class TokenizerWrapper:
201
- def __init__(self, model_path: str):
202
- model_path = str(Path(model_path))
203
- if not Path(model_path).exists():
204
- raise FileNotFoundError(model_path)
205
- self.model_path = model_path
206
- self.sp = spm.SentencePieceProcessor(model_file=model_path)
207
- self.vocab_size = int(self.sp.vocab_size())
208
- self.pad_id = self.sp.pad_id()
209
- self.unk_id = self.sp.unk_id()
210
- self.bos_id = self.sp.bos_id()
211
- self.eos_id = self.sp.eos_id()
212
- for name, val in [('pad', self.pad_id), ('unk', self.unk_id), ('bos', self.bos_id), ('eos', self.eos_id)]:
213
- if val < 0:
214
- raise ValueError(f'SentencePiece model missing <{name}>')
215
- self._special_ids = {self.pad_id, self.bos_id, self.eos_id}
216
-
217
- def encode(self, text: str, add_bos: bool = True, add_eos: bool = False) -> List[int]:
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
218
  if text is None:
219
- raise ValueError('encode() received None')
220
- if text == '':
221
  ids: List[int] = []
222
  else:
223
- ids = list(self.sp.encode(text, out_type=int))
224
  if add_bos:
225
  ids = [self.bos_id] + ids
226
  if add_eos:
227
  ids = ids + [self.eos_id]
228
  return ids
229
 
230
- def encode_large_text(
231
- self,
232
- text: str,
233
- add_bos: bool = True,
234
- add_eos: bool = False,
235
- chunk_chars: int = DEFAULT_CHUNK_CHARS,
236
- ) -> List[int]:
237
- """Backward-compatible wrapper that now delegates to the chunk iterator
238
- so encode_large_text and iter_encode_chunks share exactly one loop."""
239
- if text is None:
240
- raise ValueError('encode_large_text() received None')
241
- parts = [
242
- arr for arr in self.iter_encode_chunks(
243
- text,
244
- add_bos=add_bos,
245
- add_eos=add_eos,
246
- chunk_chars=chunk_chars,
247
- )
248
- ]
249
- if not parts:
250
- return []
251
- return np.concatenate(parts, axis=0).astype(np.int32).tolist()
252
-
253
- def iter_encode_chunks(
254
- self,
255
- text: str,
256
- add_bos: bool = True,
257
- add_eos: bool = True,
258
- chunk_chars: int = DEFAULT_CHUNK_CHARS,
259
- ):
260
- """Yield per-chunk int32 arrays while preserving the first/last chunk
261
- BOS/EOS semantics of the existing whole-file encode path.
262
-
263
- This iterator mirrors `encode_large_text()` chunk placement logic,
264
- but streams one chunk's token ids out as a NumPy array instead of
265
- materializing a Python list for the whole file.
266
- """
267
- if text is None:
268
- raise ValueError('iter_encode_chunks() received None')
269
- if text == '':
270
- ids: List[int] = []
271
- if add_bos:
272
- ids = [self.bos_id] + ids
273
- if add_eos:
274
- ids = ids + [self.eos_id]
275
- yield np.asarray(ids, dtype=np.int32)
276
- return
277
 
278
- pos = 0
279
- n = len(text)
280
- first_chunk = True
 
 
 
 
281
  while pos < n:
282
  end = min(pos + chunk_chars, n)
283
  is_eof = (end >= n)
@@ -286,26 +178,106 @@ class TokenizerWrapper:
286
  piece = window
287
  pos = end
288
  else:
289
- cut = max(window.rfind(' '), window.rfind('\n'))
290
  if cut <= 0:
291
  piece = window
292
  pos = end
293
  else:
294
  piece = window[:cut]
295
  pos += cut
296
- if not piece:
297
- if not is_eof:
298
- pos += 1
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
299
  continue
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
300
 
301
- piece_ids = self.sp.encode(piece, out_type=int)
302
- if add_bos and first_chunk:
303
- piece_ids = [self.bos_id] + piece_ids
304
- if add_eos and is_eof:
305
- piece_ids = piece_ids + [self.eos_id]
 
 
 
 
 
 
 
 
 
 
306
 
307
- yield np.asarray(piece_ids, dtype=np.int32)
308
- first_chunk = False
309
 
310
  def encode_batch(
311
  self,
@@ -314,53 +286,94 @@ class TokenizerWrapper:
314
  add_eos: bool = False,
315
  skip_errors: bool = False,
316
  ) -> List[List[int]]:
317
- out: List[List[int]] = []
318
- for i, t in enumerate(texts):
319
- try:
320
- out.append(self.encode(t, add_bos=add_bos, add_eos=add_eos))
321
- except Exception as e:
322
- if skip_errors:
 
323
  logger.warning(f"encode_batch: skipping item {i} ({e})")
324
- continue
325
- raise
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
326
  return out
327
 
328
- def decode(self, ids: Sequence[int], skip_special_tokens: bool = True) -> str:
329
- # Drop anything outside the valid piece-id range first. This is
330
- # required, not cosmetic: PyTorch's ignore_index=-100 convention for
331
- # masked label positions means `ids` is very commonly a raw labels
332
- # tensor, and sp.decode() raises IndexError on any id < 0 or
333
- # >= vocab_size instead of skipping it.
334
- ids = [int(i) for i in ids if 0 <= int(i) < self.vocab_size]
 
 
 
 
 
 
 
 
335
  if skip_special_tokens:
336
- filtered = [i for i in ids if i not in self._special_ids]
337
- else:
338
- filtered = [i for i in ids if i != self.pad_id]
339
- return self.sp.decode(filtered)
 
 
 
 
 
 
 
 
 
 
 
 
 
 
340
 
341
- def decode_batch(self, batch_ids: Sequence[Sequence[int]], skip_special_tokens: bool = True) -> List[str]:
342
- return [self.decode(ids, skip_special_tokens=skip_special_tokens) for ids in batch_ids]
343
 
344
  def save_config(self, path: str) -> None:
345
  Path(path).write_text(json.dumps({
346
- 'vocab_size': self.vocab_size,
347
- 'pad_id': self.pad_id,
348
- 'unk_id': self.unk_id,
349
- 'bos_id': self.bos_id,
350
- 'eos_id': self.eos_id,
351
- }, indent=2), encoding='utf-8')
352
 
353
  @classmethod
354
- def from_config(cls, model_path: str, config_path: Optional[str] = None) -> 'TokenizerWrapper':
355
- """Load and, if a config is given, verify special-id consistency against it."""
 
 
 
356
  tok = cls(model_path)
357
  if config_path and Path(config_path).exists():
358
- cfg = json.loads(Path(config_path).read_text(encoding='utf-8'))
359
  mismatches = {
360
  k: (cfg[k], getattr(tok, k))
361
- for k in ('vocab_size', 'pad_id', 'unk_id', 'bos_id', 'eos_id')
362
  if k in cfg and cfg[k] != getattr(tok, k)
363
  }
364
  if mismatches:
365
- raise ValueError(f'Tokenizer/config mismatch: {mismatches}')
366
  return tok
 
1
+ """
2
+ tokenizer.py — RustBPETokenizer adapter.
3
+
4
+ Drop-in replacement for the previous SentencePiece-based TokenizerWrapper.
5
+ Exposes the SAME public API so train.py, _tok_worker.py, and finetune.py
6
+ work unchanged. The only pipeline change required is that the on-disk
7
+ tokenizer artifact is now a directory containing `tokenizer.pkl` (a
8
+ pickled tiktoken.Encoding), rather than a single `tokenizer.model` file.
9
+
10
+ Set TOKENIZER_MODEL_PATH in train.py to the directory (or to the pickle
11
+ file directly — both are accepted).
12
+
13
+ API compatibility
14
+ -----------------
15
+ Public attributes:
16
+ model_path str absolute path to the .pkl (hashed into fingerprint)
17
+ vocab_size int total vocab size, specials included
18
+ bos_id int id of <|bos|>
19
+ eos_id int same as bos_id (see Notes)
20
+ pad_id int same as bos_id (see Notes)
21
+ unk_id int same as bos_id (see Notes)
22
+
23
+ Public methods:
24
+ encode(text, add_bos=True, add_eos=False) -> List[int]
25
+ encode_large_text(text, add_bos, add_eos, chunk_chars) -> List[int]
26
+ iter_encode_chunks(text, add_bos, add_eos, chunk_chars) -> Iterator[np.ndarray]
27
+ encode_batch(texts, add_bos, add_eos, skip_errors) -> List[List[int]]
28
+ decode(ids, skip_special_tokens=True) -> str
29
+ decode_batch(batch_ids, skip_special_tokens=True) -> List[str]
30
+ save_config(path)
31
+ from_config(model_path, config_path=None) -> TokenizerWrapper
32
+
33
+ Notes
34
+ -----
35
+ pad_id / eos_id / unk_id all map to bos_id because:
36
+ - tiktoken is byte-level, so <unk> is never emitted
37
+ - rustbpe's SPECIAL_TOKENS list has no dedicated <|eos>; this adapter
38
+ treats each input to encode() as one document. For iter_encode_chunks(),
39
+ one BOS is placed at the start of the file and one EOS-equivalent BOS is
40
+ placed at the end, matching the flat SP-era stream contract rather than
41
+ nanochat's per-document stream.
42
+ - finetune.py needs SOME valid id for padding; <|bos|> is the standard
43
+ choice and is filtered by decode(skip_special_tokens=True).
44
+
45
+ encode() uses tiktoken's encode_ordinary() fast path, which does NOT
46
+ interpret "<|user_start|>" etc. as special tokens. SFT rendering must
47
+ call encode_special_id() for control tokens and encode_ordinary() only
48
+ for content. This is the same contract as nanochat's trainer.
49
+ """
50
+
51
  from __future__ import annotations
52
 
53
  import json
54
  import logging
55
+ import os
56
+ import pickle
57
  from pathlib import Path
58
+ from typing import Iterator, List, Optional, Sequence, Union
59
 
60
  import numpy as np
 
 
 
61
 
62
  logger = logging.getLogger(__name__)
63
 
64
+ DEFAULT_CHUNK_CHARS = 500_000
65
+ _TIKTOKEN_THREADS = int(os.environ.get(
66
+ "TIKTOKEN_NUM_THREADS",
67
+ str(max(1, min(8, os.cpu_count() or 1))),
68
+ ))
69
+
70
+
71
+ # ---------------------------------------------------------------------------
72
+ # Helper: locate the pickle given a directory or a file path
73
+ # ---------------------------------------------------------------------------
74
+
75
+ def _resolve_pickle_path(model_path: Union[str, Path]) -> Path:
76
+ p = Path(model_path)
77
+ if p.is_dir():
78
+ candidate = p / "tokenizer.pkl"
79
+ if not candidate.exists():
80
+ raise FileNotFoundError(
81
+ f"{p} is a directory but contains no tokenizer.pkl "
82
+ f"(expected {candidate})"
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
83
  )
84
+ return candidate
85
+ if not p.exists():
86
+ raise FileNotFoundError(str(p))
87
+ return p
88
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
89
 
90
+ # ---------------------------------------------------------------------------
91
+ # TokenizerWrapper
92
+ # ---------------------------------------------------------------------------
93
 
94
  class TokenizerWrapper:
95
+ """Adapter exposing a pickled tiktoken.Encoding behind the SP-era API."""
96
+
97
+ def __init__(self, model_path: Union[str, Path]):
98
+ pickle_path = _resolve_pickle_path(model_path)
99
+ # Note: pickled tiktoken.Encoding. Only load pickles you created.
100
+ try:
101
+ with open(pickle_path, "rb") as f:
102
+ self.enc = pickle.load(f)
103
+ except Exception as exc:
104
+ raise ValueError(
105
+ f"Failed to unpickle {pickle_path}: {exc!r}. "
106
+ "This adapter expects a pickled tiktoken.Encoding "
107
+ "produced by RustBPETokenizer.save(). If you have an old "
108
+ "SentencePiece .model file, train a new tokenizer with rustbpe."
109
+ ) from exc
110
+
111
+ # Validate it looks like a tiktoken Encoding before trusting it.
112
+ for attr in ("n_vocab", "encode_ordinary", "encode_ordinary_batch",
113
+ "decode", "encode_single_token", "special_tokens_set"):
114
+ if not hasattr(self.enc, attr):
115
+ raise TypeError(
116
+ f"{pickle_path}: loaded object is not a tiktoken.Encoding "
117
+ f"(missing attribute {attr!r}). Got {type(self.enc).__name__}."
118
+ )
119
+
120
+ self.model_path = str(pickle_path)
121
+ self.vocab_size = int(self.enc.n_vocab)
122
+
123
+ # BOS is required. If the vocab lacks it, that's a hard error.
124
+ self.bos_id = int(self.enc.encode_single_token("<|bos|>"))
125
+ # EOS/PAD/UNK reuse BOS — see module docstring.
126
+ self.eos_id = self.bos_id
127
+ self.pad_id = self.bos_id
128
+ self.unk_id = self.bos_id
129
+
130
+ # Everything tiktoken labels as special — used by decode() to drop
131
+ # control tokens when skip_special_tokens=True.
132
+ self._special_ids = set()
133
+ for name in self.enc.special_tokens_set:
134
+ try:
135
+ self._special_ids.add(int(self.enc.encode_single_token(name)))
136
+ except Exception:
137
+ # If a special name isn't encodable (shouldn't happen for a
138
+ # well-formed Encoding), skip it rather than crash.
139
+ pass
140
+ # Always include BOS even if special_tokens_set was empty for some
141
+ # reason.
142
+ self._special_ids.add(self.bos_id)
143
+
144
+ # -- single-sequence encode ---------------------------------------------
145
+
146
+ def encode(
147
+ self,
148
+ text: str,
149
+ add_bos: bool = True,
150
+ add_eos: bool = False,
151
+ ) -> List[int]:
152
  if text is None:
153
+ raise ValueError("encode() received None")
154
+ if text == "":
155
  ids: List[int] = []
156
  else:
157
+ ids = self.enc.encode_ordinary(text)
158
  if add_bos:
159
  ids = [self.bos_id] + ids
160
  if add_eos:
161
  ids = ids + [self.eos_id]
162
  return ids
163
 
164
+ # -- chunked streaming --------------------------------------------------
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
165
 
166
+ @staticmethod
167
+ def _iter_chunks(text: str, chunk_chars: int) -> Iterator[str]:
168
+ """Yield whitespace-aligned chunks. Same boundaries as the SP-era
169
+ wrapper, so the resulting token stream is comparable."""
170
+ if chunk_chars <= 0:
171
+ raise ValueError(f"chunk_chars must be > 0, got {chunk_chars}")
172
+ pos, n = 0, len(text)
173
  while pos < n:
174
  end = min(pos + chunk_chars, n)
175
  is_eof = (end >= n)
 
178
  piece = window
179
  pos = end
180
  else:
181
+ cut = max(window.rfind(" "), window.rfind("\n"))
182
  if cut <= 0:
183
  piece = window
184
  pos = end
185
  else:
186
  piece = window[:cut]
187
  pos += cut
188
+ if piece:
189
+ yield piece
190
+ elif not is_eof:
191
+ pos += 1
192
+
193
+ def iter_encode_chunks(
194
+ self,
195
+ text: str,
196
+ add_bos: bool = True,
197
+ add_eos: bool = True,
198
+ chunk_chars: int = DEFAULT_CHUNK_CHARS,
199
+ batch_size: int | None = None,
200
+ ) -> Iterator[np.ndarray]:
201
+ """Yield per-chunk int32 arrays with BOS on first, EOS on last."""
202
+ if text is None:
203
+ raise ValueError("iter_encode_chunks() received None")
204
+
205
+ if text == "":
206
+ ids: List[int] = []
207
+ if add_bos:
208
+ ids.append(self.bos_id)
209
+ if add_eos:
210
+ ids.append(self.eos_id)
211
+ yield np.asarray(ids, dtype=np.int32)
212
+ return
213
+
214
+ if batch_size is None:
215
+ memory_budget_bytes = 32_000_000
216
+ chunks_per_batch = memory_budget_bytes // max(1, chunk_chars)
217
+ batch_size = max(1, min(16, chunks_per_batch))
218
+ if batch_size <= 0:
219
+ raise ValueError(f"batch_size must be > 0, got {batch_size}")
220
+ first_emitted = False
221
+ pending: List[str] = []
222
+ current: List[str] | None = None
223
+
224
+ def encode_batch(
225
+ chunks: List[str], is_last: bool, add_bos_here: bool,
226
+ ) -> Iterator[np.ndarray]:
227
+ encodings = self.enc.encode_ordinary_batch(
228
+ chunks, num_threads=_TIKTOKEN_THREADS,
229
+ )
230
+ for i, ids in enumerate(encodings):
231
+ if add_bos_here and i == 0:
232
+ ids = [self.bos_id] + ids
233
+ if add_eos and is_last and i == len(encodings) - 1:
234
+ ids = ids + [self.eos_id]
235
+ yield np.asarray(ids, dtype=np.int32)
236
+
237
+ for piece in self._iter_chunks(text, chunk_chars):
238
+ pending.append(piece)
239
+ if len(pending) < batch_size:
240
  continue
241
+ if current is not None:
242
+ yield from encode_batch(current, False, add_bos and not first_emitted)
243
+ first_emitted = True
244
+ current, pending = pending, []
245
+
246
+ if current is None:
247
+ current = pending
248
+ elif pending:
249
+ yield from encode_batch(current, False, add_bos and not first_emitted)
250
+ first_emitted = True
251
+ current = pending
252
+
253
+ if current:
254
+ yield from encode_batch(current, True, add_bos and not first_emitted)
255
+ else:
256
+ # Input was all whitespace.
257
+ ids = []
258
+ if add_bos:
259
+ ids.append(self.bos_id)
260
+ if add_eos:
261
+ ids.append(self.eos_id)
262
+ yield np.asarray(ids, dtype=np.int32)
263
 
264
+ def encode_large_text(
265
+ self,
266
+ text: str,
267
+ add_bos: bool = True,
268
+ add_eos: bool = False,
269
+ chunk_chars: int = DEFAULT_CHUNK_CHARS,
270
+ ) -> List[int]:
271
+ if text is None:
272
+ raise ValueError("encode_large_text() received None")
273
+ parts = list(self.iter_encode_chunks(
274
+ text, add_bos=add_bos, add_eos=add_eos, chunk_chars=chunk_chars,
275
+ ))
276
+ if not parts:
277
+ return []
278
+ return np.concatenate(parts, axis=0).astype(np.int32).tolist()
279
 
280
+ # -- batch ------------------------------------------------------------
 
281
 
282
  def encode_batch(
283
  self,
 
286
  add_eos: bool = False,
287
  skip_errors: bool = False,
288
  ) -> List[List[int]]:
289
+ texts = list(texts)
290
+ if skip_errors:
291
+ out: List[List[int]] = []
292
+ for i, t in enumerate(texts):
293
+ try:
294
+ out.append(self.encode(t, add_bos=add_bos, add_eos=add_eos))
295
+ except Exception as e:
296
  logger.warning(f"encode_batch: skipping item {i} ({e})")
297
+ return out
298
+
299
+ for t in texts:
300
+ if t is None:
301
+ raise ValueError("encode_batch() received None")
302
+ if not texts:
303
+ return []
304
+
305
+ encodings = self.enc.encode_ordinary_batch(
306
+ texts, num_threads=_TIKTOKEN_THREADS,
307
+ )
308
+ out = []
309
+ for ids in encodings:
310
+ if add_bos:
311
+ ids = [self.bos_id] + ids
312
+ if add_eos:
313
+ ids = ids + [self.eos_id]
314
+ out.append(ids)
315
  return out
316
 
317
+ # -- decode -----------------------------------------------------------
318
+
319
+ def decode(
320
+ self,
321
+ ids: Sequence[int],
322
+ skip_special_tokens: bool = True,
323
+ ) -> str:
324
+ # Filter out-of-range ids first: PyTorch's ignore_index=-1 convention
325
+ # for masked positions means callers frequently hand us raw label
326
+ # tensors. tiktoken.decode() raises on any id < 0 or >= vocab_size.
327
+ clean: List[int] = []
328
+ for i in ids:
329
+ iv = int(i)
330
+ if 0 <= iv < self.vocab_size:
331
+ clean.append(iv)
332
  if skip_special_tokens:
333
+ clean = [i for i in clean if i not in self._special_ids]
334
+ if not clean:
335
+ return ""
336
+ return self.enc.decode(clean)
337
+
338
+ def decode_batch(
339
+ self,
340
+ batch_ids: Sequence[Sequence[int]],
341
+ skip_special_tokens: bool = True,
342
+ ) -> List[str]:
343
+ return [self.decode(ids, skip_special_tokens=skip_special_tokens)
344
+ for ids in batch_ids]
345
+
346
+ # -- special-token helpers (extra, not in the SP wrapper) -------------
347
+
348
+ def encode_special_id(self, name: str) -> int:
349
+ """Look up the id of a named special token (e.g. '<|user_start|>')."""
350
+ return int(self.enc.encode_single_token(name))
351
 
352
+ # -- config round-trip ------------------------------------------------
 
353
 
354
  def save_config(self, path: str) -> None:
355
  Path(path).write_text(json.dumps({
356
+ "vocab_size": self.vocab_size,
357
+ "pad_id": self.pad_id,
358
+ "unk_id": self.unk_id,
359
+ "bos_id": self.bos_id,
360
+ "eos_id": self.eos_id,
361
+ }, indent=2), encoding="utf-8")
362
 
363
  @classmethod
364
+ def from_config(
365
+ cls,
366
+ model_path: str,
367
+ config_path: Optional[str] = None,
368
+ ) -> "TokenizerWrapper":
369
  tok = cls(model_path)
370
  if config_path and Path(config_path).exists():
371
+ cfg = json.loads(Path(config_path).read_text(encoding="utf-8"))
372
  mismatches = {
373
  k: (cfg[k], getattr(tok, k))
374
+ for k in ("vocab_size", "pad_id", "unk_id", "bos_id", "eos_id")
375
  if k in cfg and cfg[k] != getattr(tok, k)
376
  }
377
  if mismatches:
378
+ raise ValueError(f"Tokenizer/config mismatch: {mismatches}")
379
  return tok