Pragya / tokenizer.py
ArushBuilds's picture
Update tokenizer.py
ebd36d6
Raw History Blame Contribute Delete
14.4 kB
"""
tokenizer.py — RustBPETokenizer adapter.
Drop-in replacement for the previous SentencePiece-based TokenizerWrapper.
Exposes the SAME public API so train.py, _tok_worker.py, and finetune.py
work unchanged. The only pipeline change required is that the on-disk
tokenizer artifact is now a directory containing `tokenizer.pkl` (a
pickled tiktoken.Encoding), rather than a single `tokenizer.model` file.
Set TOKENIZER_MODEL_PATH in train.py to the directory (or to the pickle
file directly — both are accepted).
API compatibility
-----------------
Public attributes:
model_path str absolute path to the .pkl (hashed into fingerprint)
vocab_size int total vocab size, specials included
bos_id int id of <|bos|>
eos_id int same as bos_id (see Notes)
pad_id int same as bos_id (see Notes)
unk_id int same as bos_id (see Notes)
Public methods:
encode(text, add_bos=True, add_eos=False) -> List[int]
encode_large_text(text, add_bos, add_eos, chunk_chars) -> List[int]
iter_encode_chunks(text, add_bos, add_eos, chunk_chars) -> Iterator[np.ndarray]
encode_batch(texts, add_bos, add_eos, skip_errors) -> List[List[int]]
decode(ids, skip_special_tokens=True) -> str
decode_batch(batch_ids, skip_special_tokens=True) -> List[str]
save_config(path)
from_config(model_path, config_path=None) -> TokenizerWrapper
Notes
-----
pad_id / eos_id / unk_id all map to bos_id because:
- tiktoken is byte-level, so <unk> is never emitted
- rustbpe's SPECIAL_TOKENS list has no dedicated <|eos>; this adapter
treats each input to encode() as one document. For iter_encode_chunks(),
one BOS is placed at the start of the file and one EOS-equivalent BOS is
placed at the end, matching the flat SP-era stream contract rather than
nanochat's per-document stream.
- finetune.py needs SOME valid id for padding; <|bos|> is the standard
choice and is filtered by decode(skip_special_tokens=True).
encode() uses tiktoken's encode_ordinary() fast path, which does NOT
interpret "<|user_start|>" etc. as special tokens. SFT rendering must
call encode_special_id() for control tokens and encode_ordinary() only
for content. This is the same contract as nanochat's trainer.
"""
from __future__ import annotations
import json
import logging
import os
import pickle
from pathlib import Path
from typing import Iterator, List, Optional, Sequence, Union
import numpy as np
logger = logging.getLogger(__name__)
DEFAULT_CHUNK_CHARS = 500_000
_TIKTOKEN_THREADS = int(os.environ.get(
"TIKTOKEN_NUM_THREADS",
str(max(1, min(8, os.cpu_count() or 1))),
))
# ---------------------------------------------------------------------------
# Helper: locate the pickle given a directory or a file path
# ---------------------------------------------------------------------------
def _resolve_pickle_path(model_path: Union[str, Path]) -> Path:
p = Path(model_path)
if p.is_dir():
candidate = p / "tokenizer.pkl"
if not candidate.exists():
raise FileNotFoundError(
f"{p} is a directory but contains no tokenizer.pkl "
f"(expected {candidate})"
)
return candidate
if not p.exists():
raise FileNotFoundError(str(p))
return p
# ---------------------------------------------------------------------------
# TokenizerWrapper
# ---------------------------------------------------------------------------
class TokenizerWrapper:
"""Adapter exposing a pickled tiktoken.Encoding behind the SP-era API."""
def __init__(self, model_path: Union[str, Path]):
pickle_path = _resolve_pickle_path(model_path)
# Note: pickled tiktoken.Encoding. Only load pickles you created.
try:
with open(pickle_path, "rb") as f:
self.enc = pickle.load(f)
except Exception as exc:
raise ValueError(
f"Failed to unpickle {pickle_path}: {exc!r}. "
"This adapter expects a pickled tiktoken.Encoding "
"produced by RustBPETokenizer.save(). If you have an old "
"SentencePiece .model file, train a new tokenizer with rustbpe."
) from exc
# Validate it looks like a tiktoken Encoding before trusting it.
for attr in ("n_vocab", "encode_ordinary", "encode_ordinary_batch",
"decode", "encode_single_token", "special_tokens_set"):
if not hasattr(self.enc, attr):
raise TypeError(
f"{pickle_path}: loaded object is not a tiktoken.Encoding "
f"(missing attribute {attr!r}). Got {type(self.enc).__name__}."
)
self.model_path = str(pickle_path)
self.vocab_size = int(self.enc.n_vocab)
# BOS is required. If the vocab lacks it, that's a hard error.
self.bos_id = int(self.enc.encode_single_token("<|bos|>"))
# EOS/PAD/UNK reuse BOS — see module docstring.
self.eos_id = self.bos_id
self.pad_id = self.bos_id
self.unk_id = self.bos_id
# Everything tiktoken labels as special — used by decode() to drop
# control tokens when skip_special_tokens=True.
self._special_ids = set()
for name in self.enc.special_tokens_set:
try:
self._special_ids.add(int(self.enc.encode_single_token(name)))
except Exception:
# If a special name isn't encodable (shouldn't happen for a
# well-formed Encoding), skip it rather than crash.
pass
# Always include BOS even if special_tokens_set was empty for some
# reason.
self._special_ids.add(self.bos_id)
# -- single-sequence encode ---------------------------------------------
def encode(
self,
text: str,
add_bos: bool = True,
add_eos: bool = False,
) -> List[int]:
if text is None:
raise ValueError("encode() received None")
if text == "":
ids: List[int] = []
else:
ids = self.enc.encode_ordinary(text)
if add_bos:
ids = [self.bos_id] + ids
if add_eos:
ids = ids + [self.eos_id]
return ids
# -- chunked streaming --------------------------------------------------
@staticmethod
def _iter_chunks(text: str, chunk_chars: int) -> Iterator[str]:
"""Yield whitespace-aligned chunks. Same boundaries as the SP-era
wrapper, so the resulting token stream is comparable."""
if chunk_chars <= 0:
raise ValueError(f"chunk_chars must be > 0, got {chunk_chars}")
pos, n = 0, len(text)
while pos < n:
end = min(pos + chunk_chars, n)
is_eof = (end >= n)
window = text[pos:end]
if is_eof:
piece = window
pos = end
else:
cut = max(window.rfind(" "), window.rfind("\n"))
if cut <= 0:
piece = window
pos = end
else:
piece = window[:cut]
pos += cut
if piece:
yield piece
elif not is_eof:
pos += 1
def iter_encode_chunks(
self,
text: str,
add_bos: bool = True,
add_eos: bool = True,
chunk_chars: int = DEFAULT_CHUNK_CHARS,
batch_size: int | None = None,
) -> Iterator[np.ndarray]:
"""Yield per-chunk int32 arrays with BOS on first, EOS on last."""
if text is None:
raise ValueError("iter_encode_chunks() received None")
if text == "":
ids: List[int] = []
if add_bos:
ids.append(self.bos_id)
if add_eos:
ids.append(self.eos_id)
yield np.asarray(ids, dtype=np.int32)
return
if batch_size is None:
memory_budget_bytes = 32_000_000
chunks_per_batch = memory_budget_bytes // max(1, chunk_chars)
batch_size = max(1, min(16, chunks_per_batch))
if batch_size <= 0:
raise ValueError(f"batch_size must be > 0, got {batch_size}")
first_emitted = False
pending: List[str] = []
current: List[str] | None = None
def encode_batch(
chunks: List[str], is_last: bool, add_bos_here: bool,
) -> Iterator[np.ndarray]:
encodings = self.enc.encode_ordinary_batch(
chunks, num_threads=_TIKTOKEN_THREADS,
)
for i, ids in enumerate(encodings):
if add_bos_here and i == 0:
ids = [self.bos_id] + ids
if add_eos and is_last and i == len(encodings) - 1:
ids = ids + [self.eos_id]
yield np.asarray(ids, dtype=np.int32)
for piece in self._iter_chunks(text, chunk_chars):
pending.append(piece)
if len(pending) < batch_size:
continue
if current is not None:
yield from encode_batch(current, False, add_bos and not first_emitted)
first_emitted = True
current, pending = pending, []
if current is None:
current = pending
elif pending:
yield from encode_batch(current, False, add_bos and not first_emitted)
first_emitted = True
current = pending
if current:
yield from encode_batch(current, True, add_bos and not first_emitted)
else:
# Input was all whitespace.
ids = []
if add_bos:
ids.append(self.bos_id)
if add_eos:
ids.append(self.eos_id)
yield np.asarray(ids, dtype=np.int32)
def encode_large_text(
self,
text: str,
add_bos: bool = True,
add_eos: bool = False,
chunk_chars: int = DEFAULT_CHUNK_CHARS,
) -> List[int]:
if text is None:
raise ValueError("encode_large_text() received None")
parts = list(self.iter_encode_chunks(
text, add_bos=add_bos, add_eos=add_eos, chunk_chars=chunk_chars,
))
if not parts:
return []
return np.concatenate(parts, axis=0).astype(np.int32).tolist()
# -- batch ------------------------------------------------------------
def encode_batch(
self,
texts: Sequence[str],
add_bos: bool = True,
add_eos: bool = False,
skip_errors: bool = False,
) -> List[List[int]]:
texts = list(texts)
if skip_errors:
out: List[List[int]] = []
for i, t in enumerate(texts):
try:
out.append(self.encode(t, add_bos=add_bos, add_eos=add_eos))
except Exception as e:
logger.warning(f"encode_batch: skipping item {i} ({e})")
return out
for t in texts:
if t is None:
raise ValueError("encode_batch() received None")
if not texts:
return []
encodings = self.enc.encode_ordinary_batch(
texts, num_threads=_TIKTOKEN_THREADS,
)
out = []
for ids in encodings:
if add_bos:
ids = [self.bos_id] + ids
if add_eos:
ids = ids + [self.eos_id]
out.append(ids)
return out
# -- decode -----------------------------------------------------------
def decode(
self,
ids: Sequence[int],
skip_special_tokens: bool = True,
) -> str:
# Filter out-of-range ids first: PyTorch's ignore_index=-1 convention
# for masked positions means callers frequently hand us raw label
# tensors. tiktoken.decode() raises on any id < 0 or >= vocab_size.
clean: List[int] = []
for i in ids:
iv = int(i)
if 0 <= iv < self.vocab_size:
clean.append(iv)
if skip_special_tokens:
clean = [i for i in clean if i not in self._special_ids]
if not clean:
return ""
return self.enc.decode(clean)
def decode_batch(
self,
batch_ids: Sequence[Sequence[int]],
skip_special_tokens: bool = True,
) -> List[str]:
return [self.decode(ids, skip_special_tokens=skip_special_tokens)
for ids in batch_ids]
# -- special-token helpers (extra, not in the SP wrapper) -------------
def encode_special_id(self, name: str) -> int:
"""Look up the id of a named special token (e.g. '<|user_start|>')."""
return int(self.enc.encode_single_token(name))
# -- config round-trip ------------------------------------------------
def save_config(self, path: str) -> None:
Path(path).write_text(json.dumps({
"vocab_size": self.vocab_size,
"pad_id": self.pad_id,
"unk_id": self.unk_id,
"bos_id": self.bos_id,
"eos_id": self.eos_id,
}, indent=2), encoding="utf-8")
@classmethod
def from_config(
cls,
model_path: str,
config_path: Optional[str] = None,
) -> "TokenizerWrapper":
tok = cls(model_path)
if config_path and Path(config_path).exists():
cfg = json.loads(Path(config_path).read_text(encoding="utf-8"))
mismatches = {
k: (cfg[k], getattr(tok, k))
for k in ("vocab_size", "pad_id", "unk_id", "bos_id", "eos_id")
if k in cfg and cfg[k] != getattr(tok, k)
}
if mismatches:
raise ValueError(f"Tokenizer/config mismatch: {mismatches}")
return tok