""" GCTokenizer-v1 reference implementation. Deterministic longest-match segmentation over UTF-8 bytes with universal single-byte fallback tokens of the form <0xXX>. """ from __future__ import annotations import json import re from pathlib import Path from typing import Dict, Iterable, List, Optional, Sequence, Tuple _BYTE_TOKEN_RE = re.compile(r"^<0x([0-9A-Fa-f]{2})>$") class _TrieNode: __slots__ = ("children", "token_id") def __init__(self) -> None: self.children: Dict[int, "_TrieNode"] = {} self.token_id: Optional[int] = None class GCTokenizer: """Lossless longest-match tokenizer with byte fallback.""" def __init__(self, vocab: Dict[str, int], *, name: str = "GCT") -> None: self.name = name self.token_to_id: Dict[str, int] = dict(vocab) self.id_to_token: Dict[int, str] = {i: t for t, i in vocab.items()} # id -> raw bytes emitted on decode / matched on encode self.id_to_bytes: Dict[int, bytes] = {} self.byte_id: List[Optional[int]] = [None] * 256 root = _TrieNode() for token, tid in vocab.items(): m = _BYTE_TOKEN_RE.match(token) if m: b = int(m.group(1), 16) bb = bytes([b]) self.id_to_bytes[tid] = bb self.byte_id[b] = tid else: bb = token.encode("utf-8") self.id_to_bytes[tid] = bb node = root for byte in bb: nxt = node.children.get(byte) if nxt is None: nxt = _TrieNode() node.children[byte] = nxt node = nxt # Prefer keeping an existing terminal if somehow duplicated; # vocab ids are unique so this just records the token. node.token_id = tid self._root = root missing = [i for i, tid in enumerate(self.byte_id) if tid is None] if missing: raise ValueError( f"{name}: byte fallback incomplete, missing bytes {missing[:8]}..." ) @classmethod def from_dir(cls, path: str | Path) -> "GCTokenizer": path = Path(path) data = json.loads((path / "tokenizer.json").read_text(encoding="utf-8")) if data.get("type") not in (None, "GCT"): # Still accept; GCT is the expected type. pass vocab = data["vocab"] name = path.name return cls(vocab, name=name) # ------------------------------------------------------------------ # Core API # ------------------------------------------------------------------ def encode_bytes(self, data: bytes) -> List[int]: """Greedy longest-match encode over raw bytes.""" root = self._root ids: List[int] = [] n = len(data) i = 0 while i < n: node = root last_id: Optional[int] = None last_j = i j = i while j < n: nxt = node.children.get(data[j]) if nxt is None: break node = nxt j += 1 if node.token_id is not None: last_id = node.token_id last_j = j if last_id is not None: ids.append(last_id) i = last_j else: # Single-byte fallback (always present). ids.append(self.byte_id[data[i]]) # type: ignore[arg-type] i += 1 return ids def decode_bytes(self, ids: Sequence[int]) -> bytes: parts = self.id_to_bytes out = bytearray() for tid in ids: b = parts.get(tid) if b is None: raise KeyError(f"unknown token id: {tid}") out.extend(b) return bytes(out) def encode(self, text: str) -> List[int]: return self.encode_bytes(text.encode("utf-8")) def decode(self, ids: Sequence[int]) -> str: return self.decode_bytes(ids).decode("utf-8") def tokenize(self, text: str) -> List[str]: return [self.id_to_token[i] for i in self.encode(text)] def roundtrip_bytes(self, data: bytes) -> bool: return self.decode_bytes(self.encode_bytes(data)) == data def __len__(self) -> int: return len(self.token_to_id) def load_all(root: str | Path = "GCTokenizer-v1") -> Dict[str, GCTokenizer]: root = Path(root) return {tier: GCTokenizer.from_dir(root / tier) for tier in ("S", "M", "L", "XL")}