Spaces:
Sleeping
Sleeping
Download tokenizer.py from ArushBuilds/Pragya: direct link, hf CLI and curl.
- Browser
- Download file 14.4 kB
-
https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/main/tokenizer.py
- Command line
-
hf download hf://spaces/ArushBuilds/Pragya/tokenizer.py
-
curl -L -o tokenizer.py https://huggingface.co/spaces/ArushBuilds/Pragya/resolve/main/tokenizer.py
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 -------------------------------------------------- | |
| 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") | |
| 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 |