GCTokenizer-v1 / gct_tokenizer.py
wop's picture
Upload gct_tokenizer.py
acd6bb5 verified
Raw
History Blame Contribute Delete
4.59 kB
"""
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")}