| """ |
| 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()} |
|
|
| |
| 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 |
| |
| |
| 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"): |
| |
| pass |
| vocab = data["vocab"] |
| name = path.name |
| return cls(vocab, name=name) |
|
|
| |
| |
| |
| 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: |
| |
| ids.append(self.byte_id[data[i]]) |
| 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")} |
|
|