boopit-1 / tokenizer.py
ProCreations's picture
Upload Boopit 1.58-bit packed weights
2cb0b71 verified
Raw
History Blame Contribute Delete
4.67 kB
"""Byte-level BPE compatible with babble / booper-pretrain tokenizer.json."""
from __future__ import annotations
import json
import re
from pathlib import Path
_CHUNK_RE = re.compile(r"\s+|\S+")
def _merge_ids(ids: list[int], a: int, b: int, new_id: int) -> list[int]:
if len(ids) < 2:
return ids
out: list[int] = []
i = 0
n = len(ids)
while i < n:
if i + 1 < n and ids[i] == a and ids[i + 1] == b:
out.append(new_id)
i += 2
else:
out.append(ids[i])
i += 1
return out
class BPETokenizer:
def __init__(self, merges: list[tuple[int, int, int]]) -> None:
self.merges = merges
vocab: dict[int, bytes] = {i: bytes([i]) for i in range(256)}
for a, b, new_id in merges:
vocab[new_id] = vocab[a] + vocab[b]
self.vocab = vocab
self._ranks = {(a, b): i for i, (a, b, _) in enumerate(merges)}
self._pair_to_id = {(a, b): new_id for a, b, new_id in merges}
base = 256 + len(merges)
self.pad, self.bos, self.sep, self.eos = base, base + 1, base + 2, base + 3
@property
def vocab_size(self) -> int:
return 256 + len(self.merges) + 4
def _encode_chunk(self, chunk: str) -> list[int]:
ids = list(chunk.encode("utf-8"))
ranks = self._ranks
pair_to_id = self._pair_to_id
while len(ids) >= 2:
best_rank = None
best_pair = None
for a, b in zip(ids, ids[1:]):
r = ranks.get((a, b))
if r is not None and (best_rank is None or r < best_rank):
best_rank = r
best_pair = (a, b)
if best_pair is None:
break
ids = _merge_ids(ids, best_pair[0], best_pair[1], pair_to_id[best_pair])
return ids
def _build_fast(self):
try:
from tokenizers import Tokenizer
from tokenizers import models as tokmodels
except Exception:
self._fast = None
return
id_to_tok = {i: bytes([i]).decode("latin-1") for i in range(256)}
vocab = {s: i for i, s in id_to_tok.items()}
hf_merges: list[tuple[str, str]] = []
for a, b, nid in self.merges:
sa, sb = id_to_tok[a], id_to_tok[b]
merged = sa + sb
id_to_tok[nid] = merged
vocab[merged] = nid
hf_merges.append((sa, sb))
fast = Tokenizer(tokmodels.BPE(vocab, hf_merges, fuse_unk=False))
self._fast = fast
def _ensure_fast(self) -> None:
if getattr(self, "_fast", None) is None and not hasattr(self, "_fast_tried"):
self._fast_tried = True
self._build_fast()
def encode(self, text: str) -> list[int]:
self._ensure_fast()
if getattr(self, "_fast", None) is not None:
ids: list[int] = []
for chunk in _CHUNK_RE.findall(text):
raw = chunk.encode("utf-8").decode("latin-1")
ids.extend(self._fast.encode(raw).ids)
return ids
ids = []
for chunk in _CHUNK_RE.findall(text):
ids.extend(self._encode_chunk(chunk))
return ids
def encode_docs(self, texts: list[str]) -> list[int]:
"""Encode many docs and join with eos. Uses tokenizers encode_batch."""
self._ensure_fast()
chunks: list[str] = []
lens: list[int] = []
for text in texts:
cs = _CHUNK_RE.findall(text)
lens.append(len(cs))
chunks.extend(c.encode("utf-8").decode("latin-1") for c in cs)
out: list[int] = []
if self._fast is not None and chunks:
encs = self._fast.encode_batch(chunks)
i = 0
for n in lens:
for _ in range(n):
out.extend(encs[i].ids)
i += 1
out.append(self.eos)
return out
for text in texts:
out.extend(self.encode(text))
out.append(self.eos)
return out
def decode(self, ids: list[int]) -> str:
raw = bytearray()
for i in ids:
piece = self.vocab.get(i)
if piece is not None:
raw.extend(piece)
return bytes(raw).decode("utf-8", errors="replace")
def to_json(self, path: Path) -> None:
path.write_text(json.dumps({"merges": [list(m) for m in self.merges]}))
@classmethod
def from_json(cls, path: Path | str) -> "BPETokenizer":
raw = json.loads(Path(path).read_text())
return cls([tuple(m) for m in raw["merges"]])