Spaces:
Sleeping
Sleeping
Arush kumar commited on
Commit ·
71f1a2a
1
Parent(s): cbc2499
Update tokenizer.py
Browse files- 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 |
-
|
| 9 |
-
from tokenizers.models import BPE
|
| 10 |
-
from tokenizers.processors import TemplateProcessing
|
| 11 |
-
import tiktoken
|
| 12 |
|
| 13 |
logger = logging.getLogger(__name__)
|
| 14 |
|
| 15 |
-
#
|
| 16 |
-
#
|
| 17 |
-
#
|
| 18 |
-
#
|
| 19 |
-
|
| 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 =
|
| 34 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 35 |
) -> str:
|
| 36 |
"""
|
| 37 |
-
Train a
|
| 38 |
-
|
| 39 |
-
|
| 40 |
-
|
| 41 |
-
Why
|
| 42 |
-
|
| 43 |
-
|
| 44 |
-
|
| 45 |
-
|
| 46 |
-
|
| 47 |
-
|
| 48 |
-
|
| 49 |
-
|
| 50 |
-
|
| 51 |
-
|
| 52 |
-
|
| 53 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 64 |
-
|
| 65 |
-
|
| 66 |
-
|
| 67 |
-
|
| 68 |
-
|
| 69 |
-
|
| 70 |
-
|
| 71 |
-
|
| 72 |
-
|
| 73 |
-
|
| 74 |
-
|
| 75 |
-
|
| 76 |
-
|
| 77 |
-
|
| 78 |
-
|
| 79 |
-
|
| 80 |
-
|
| 81 |
-
|
| 82 |
-
|
| 83 |
-
|
| 84 |
-
("<eos>", tokenizer.token_to_id("<eos>")),
|
| 85 |
-
],
|
| 86 |
)
|
| 87 |
|
| 88 |
-
|
| 89 |
-
|
| 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 |
-
|
| 106 |
-
|
| 107 |
-
|
| 108 |
-
|
| 109 |
-
+ list(range(ord("¡"), ord("¬") + 1))
|
| 110 |
-
+ list(range(ord("®"), ord("ÿ") + 1))
|
| 111 |
)
|
| 112 |
-
|
| 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
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 152 |
"""
|
| 153 |
-
|
| 154 |
-
|
| 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 |
-
|
| 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 |
-
|
| 242 |
-
|
| 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 |
-
|
| 250 |
-
|
| 251 |
-
|
| 252 |
-
|
| 253 |
-
|
| 254 |
-
|
| 255 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 256 |
raise ValueError('Validation encode produced empty output')
|
| 257 |
-
decoded =
|
| 258 |
if not decoded.strip():
|
| 259 |
raise ValueError('Validation round-trip produced empty decode')
|
| 260 |
|
| 261 |
-
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
|
| 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 |
-
|
| 273 |
-
|
| 274 |
-
|
| 275 |
-
|
| 276 |
-
|
| 277 |
-
|
| 278 |
-
|
| 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
|
| 287 |
-
raise ValueError(f'
|
| 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 |
-
|
| 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.
|
| 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'))
|