""" 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 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