Spaces:
Running
Running
Commit ·
ebd36d6
1
Parent(s): a3032cf
Update tokenizer.py
Browse files- 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
|
| 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 |
-
|
| 16 |
-
|
| 17 |
-
|
| 18 |
-
|
| 19 |
-
|
| 20 |
-
|
| 21 |
-
|
| 22 |
-
|
| 23 |
-
|
| 24 |
-
|
| 25 |
-
|
| 26 |
-
|
| 27 |
-
|
| 28 |
-
|
| 29 |
-
|
| 30 |
-
|
| 31 |
-
|
| 32 |
-
|
| 33 |
-
|
| 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 |
-
|
| 202 |
-
|
| 203 |
-
|
| 204 |
-
|
| 205 |
-
|
| 206 |
-
|
| 207 |
-
|
| 208 |
-
|
| 209 |
-
|
| 210 |
-
|
| 211 |
-
|
| 212 |
-
|
| 213 |
-
|
| 214 |
-
|
| 215 |
-
|
| 216 |
-
|
| 217 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 218 |
if text is None:
|
| 219 |
-
raise ValueError(
|
| 220 |
-
if text ==
|
| 221 |
ids: List[int] = []
|
| 222 |
else:
|
| 223 |
-
ids =
|
| 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 |
-
|
| 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 |
-
|
| 279 |
-
|
| 280 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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(
|
| 290 |
if cut <= 0:
|
| 291 |
piece = window
|
| 292 |
pos = end
|
| 293 |
else:
|
| 294 |
piece = window[:cut]
|
| 295 |
pos += cut
|
| 296 |
-
if
|
| 297 |
-
|
| 298 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 299 |
continue
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 300 |
|
| 301 |
-
|
| 302 |
-
|
| 303 |
-
|
| 304 |
-
|
| 305 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 306 |
|
| 307 |
-
|
| 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 |
-
|
| 318 |
-
|
| 319 |
-
|
| 320 |
-
|
| 321 |
-
|
| 322 |
-
|
|
|
|
| 323 |
logger.warning(f"encode_batch: skipping item {i} ({e})")
|
| 324 |
-
|
| 325 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 326 |
return out
|
| 327 |
|
| 328 |
-
|
| 329 |
-
|
| 330 |
-
|
| 331 |
-
|
| 332 |
-
|
| 333 |
-
|
| 334 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 335 |
if skip_special_tokens:
|
| 336 |
-
|
| 337 |
-
|
| 338 |
-
|
| 339 |
-
return self.
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 340 |
|
| 341 |
-
|
| 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 |
-
|
| 347 |
-
|
| 348 |
-
|
| 349 |
-
|
| 350 |
-
|
| 351 |
-
}, indent=2), encoding=
|
| 352 |
|
| 353 |
@classmethod
|
| 354 |
-
def from_config(
|
| 355 |
-
|
|
|
|
|
|
|
|
|
|
| 356 |
tok = cls(model_path)
|
| 357 |
if config_path and Path(config_path).exists():
|
| 358 |
-
cfg = json.loads(Path(config_path).read_text(encoding=
|
| 359 |
mismatches = {
|
| 360 |
k: (cfg[k], getattr(tok, k))
|
| 361 |
-
for k in (
|
| 362 |
if k in cfg and cfg[k] != getattr(tok, k)
|
| 363 |
}
|
| 364 |
if mismatches:
|
| 365 |
-
raise ValueError(f
|
| 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
|