File size: 30,679 Bytes
efdf5ed | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 237 238 239 240 241 242 243 244 245 246 247 248 249 250 251 252 253 254 255 256 257 258 259 260 261 262 263 264 265 266 267 268 269 270 271 272 273 274 275 276 277 278 279 280 281 282 283 284 285 286 287 288 289 290 291 292 293 294 295 296 297 298 299 300 301 302 303 304 305 306 307 308 309 310 311 312 313 314 315 316 317 318 319 320 321 322 323 324 325 326 327 328 329 330 331 332 333 334 335 336 337 338 339 340 341 342 343 344 345 346 347 348 349 350 351 352 353 354 355 356 357 358 359 360 361 362 363 364 365 366 367 368 369 370 371 372 373 374 375 376 377 378 379 380 381 382 383 384 385 386 387 388 389 390 391 392 393 394 395 396 397 398 399 400 401 402 403 404 405 406 407 408 409 410 411 412 413 414 415 416 417 418 419 420 421 422 423 424 425 426 427 428 429 430 431 432 433 434 435 436 437 438 439 440 441 442 443 444 445 446 447 448 449 450 451 452 453 454 455 456 457 458 459 460 461 462 463 464 465 466 467 468 469 470 471 472 473 474 475 476 477 478 479 480 481 482 483 484 485 486 487 488 489 490 491 492 493 494 495 496 497 498 499 500 501 502 503 504 505 506 507 508 509 510 511 512 513 514 515 516 517 518 519 520 521 522 523 524 525 526 527 528 529 530 531 532 533 534 535 536 537 538 539 540 541 542 543 544 545 546 547 548 549 550 551 552 553 554 555 556 557 558 559 560 561 562 563 564 565 566 567 568 569 570 571 572 573 574 575 576 577 578 579 580 581 582 583 584 585 586 587 588 589 590 591 592 593 594 595 596 597 598 599 600 601 602 603 604 605 606 607 608 609 610 611 612 613 614 615 616 617 618 619 620 621 622 623 624 625 626 627 628 629 630 631 632 633 634 635 636 637 638 639 640 641 642 643 644 645 646 647 648 649 650 651 652 653 654 655 656 657 658 659 660 661 662 663 664 665 666 667 668 669 670 671 672 673 674 675 676 677 678 679 680 681 682 683 684 685 686 687 688 689 690 691 692 693 694 695 696 697 698 699 700 701 702 703 704 705 706 707 708 709 710 711 712 713 714 715 716 717 718 719 720 721 722 723 724 725 726 727 728 729 730 731 732 733 734 735 736 737 738 739 740 741 742 743 744 745 746 747 748 749 750 751 752 753 754 755 756 757 758 759 760 761 762 763 764 765 766 767 768 769 770 771 772 773 774 775 776 777 778 779 780 781 782 783 784 785 786 787 788 789 790 791 792 793 794 795 796 797 798 799 800 801 802 803 804 805 806 807 808 809 810 811 812 813 814 815 816 817 818 819 820 821 822 823 824 825 826 827 828 829 830 831 | #!/usr/bin/env python3
# SPDX-License-Identifier: MIT
# Copyright (c) 2026 Hamid Wakili <hamid@ideployed.com>
"""A compact Byte-Pair Encoding tokenizer implementation.
This module implements a byte-level Byte-Pair Encoding (BPE) tokenizer using
only the Python standard library. It learns every merge and vocabulary entry
from the supplied corpus; it does not load a pretrained tokenizer vocabulary.
Algorithmic background
----------------------
BPE was introduced for neural machine translation by Sennrich, Haddow, and
Birch in "Neural Machine Translation of Rare Words with Subword Units" (2016).
The training loop repeatedly finds the most frequent adjacent symbol pair in a
corpus and replaces that pair with a newly created merged symbol. The result is
a compact subword vocabulary learned entirely from the supplied corpus.
Design choice: byte-level BPE
-----------------------------
This implementation operates on UTF-8 bytes instead of Unicode code points.
German text contains umlauts, sharp-s, names, punctuation, and possibly mixed
foreign terms. Byte-level BPE can represent any valid Unicode input because the
base vocabulary covers all 256 byte values. That means encoding does not need an
imported character vocabulary and rarely needs ``<unk>``. The end-of-word marker
``</w>`` is represented as a special atomic symbol after every whitespace-
separated word, matching the classic BPE word-boundary convention.
The implementation prioritizes correctness, transparency, and inspectable
artifacts. The original ``train`` method keeps the simple full-recount BPE loop
for auditability and small corpora. The ``train_fast`` method uses weighted
unique words, live pair counts, an inverted pair-to-word index, and a
lazy-invalidated heap so larger tokenizer runs finish without relying
on external tokenizer packages.
Implementation note
-------------------
This implementation was written for this project using the BPE algorithm
described above. It does not import merge rules, vocabulary entries, or
tokenizer artifacts from an existing model.
"""
from __future__ import annotations
import json
import heapq
import multiprocessing as mp
import re
import sys
import tempfile
import time
from collections import Counter
from pathlib import Path
from typing import Dict, Iterable, List, MutableMapping, Sequence, Tuple
Symbol = str
Pair = Tuple[Symbol, Symbol]
Word = Tuple[Symbol, ...]
AtomWord = Tuple[int, ...]
def _count_words_chunk(texts: Sequence[str]) -> Counter[str]:
"""Count whitespace-normalized words in a chunk of texts.
This top-level helper is intentionally separate from ``BPE_Tokenizer`` so it
can be used by ``multiprocessing.Pool`` on platforms that require picklable
worker functions.
"""
counts: Counter[str] = Counter()
for text in texts:
normalized = re.sub(r"\s+", " ", str(text).strip())
if normalized:
counts.update(normalized.split(" "))
return counts
class BPE_Tokenizer:
"""Train, save, load, encode, and decode a byte-level BPE tokenizer.
The tokenizer learns merge rules from a UTF-8 corpus. Its initial atomic
vocabulary has a fixed, deterministic base layout:
* ID 0: ``<unk>``
* ID 1: ``<pad>``
* ID 2: ``<s>``
* ID 3: ``</s>``
* ID 4: ``</w>``
* IDs 5-260: byte values 0-255
* IDs 261 and above: learned BPE merge symbols
Tokens are stored internally as JSON-safe symbol strings. Byte 65 is stored
as ``"65"``. A merged symbol is stored as a space-separated sequence of
atomic units such as ``"76 101 105"`` for the bytes of ``"Lei"``. A token
that includes the word boundary marker may end in ``"256"``, the internal
atom value for ``</w>``.
Parameters
----------
special_tokens:
Optional mapping of token string to integer ID. If omitted, sensible
defaults are used. The ``</w>`` token must be present because it is part
of the BPE training data.
Attributes
----------
vocab:
Mapping from symbol string to token ID.
id_to_symbol:
Reverse mapping from token ID to symbol string.
merges:
Ordered list of merge pairs. Encoding applies these rules in order.
"""
DEFAULT_SPECIAL_TOKENS: Dict[str, int] = {
"<unk>": 0,
"<pad>": 1,
"<s>": 2,
"</s>": 3,
"</w>": 4,
}
EOW_ATOM = 256
EOW_SYMBOL = "</w>"
def __init__(self, special_tokens: MutableMapping[str, int] | None = None) -> None:
self.special_tokens: Dict[str, int] = dict(
special_tokens or self.DEFAULT_SPECIAL_TOKENS
)
self._validate_special_tokens(self.special_tokens)
self.vocab: Dict[Symbol, int] = {}
self.id_to_symbol: Dict[int, Symbol] = {}
self.merges: List[Pair] = []
self.merge_ranks: Dict[Pair, int] = {}
self._initialize_base_vocabulary()
@classmethod
def train_from_file(
cls,
path: str | Path,
num_merges: int,
special_tokens: MutableMapping[str, int] | None = None,
min_pair_frequency: int = 2,
) -> "BPE_Tokenizer":
"""Create and train a tokenizer from a UTF-8 text file.
The file is read as UTF-8, whitespace is normalized, words are split on
single spaces, and ``</w>`` is appended to every word before pair counts
are computed.
Parameters
----------
path:
Raw UTF-8 corpus file.
num_merges:
Maximum number of BPE merge operations to learn.
special_tokens:
Optional special token mapping. By default IDs 0-4 are reserved for
``<unk>``, ``<pad>``, ``<s>``, ``</s>``, and ``</w>``.
min_pair_frequency:
Stop training when the best pair frequency is below this threshold.
The default of 2 avoids creating tokens that appear only once.
Returns
-------
BPE_Tokenizer
A trained tokenizer instance.
"""
tokenizer = cls(special_tokens=special_tokens)
text = Path(path).read_text(encoding="utf-8")
tokenizer.train(text, num_merges=num_merges, min_pair_frequency=min_pair_frequency)
return tokenizer
def train(
self,
text: str,
num_merges: int,
min_pair_frequency: int = 2,
) -> None:
"""Learn BPE merge rules from raw text.
This method implements the standard BPE training loop:
1. Convert every normalized word into byte symbols plus ``</w>``.
2. Count adjacent symbol pairs across the corpus.
3. Merge the most frequent pair into a new symbol.
4. Replace all occurrences of that pair in the corpus.
5. Repeat until ``num_merges`` is reached or no useful pair remains.
The corpus is compressed into a frequency dictionary of unique word
symbol sequences before training. Pair counts are rebuilt after every
merge. This is simpler than an incremental priority queue and is easier
to audit, at the cost of extra training time on very large corpora.
"""
if num_merges < 0:
raise ValueError("num_merges must be non-negative")
if min_pair_frequency < 1:
raise ValueError("min_pair_frequency must be at least 1")
corpus = self._build_training_corpus(text)
for _ in range(num_merges):
pair_counts = self._count_pairs(corpus)
if not pair_counts:
break
best_pair, best_count = self._best_pair(pair_counts)
if best_count < min_pair_frequency:
break
new_symbol = self._merge_symbol(best_pair)
if new_symbol not in self.vocab:
self._add_symbol(new_symbol)
self.merges.append(best_pair)
self.merge_ranks[best_pair] = len(self.merges) - 1
corpus = self._replace_pair_in_corpus(corpus, best_pair, new_symbol)
def train_fast(
self,
texts: str | Sequence[str],
num_merges: int,
min_pair_frequency: int = 2,
n_workers: int | None = None,
chunk_size: int = 512,
verbose: bool = True,
) -> None:
"""Learn BPE merge rules with an incremental heap-based trainer.
The public output is identical in format to :meth:`train`: ``vocab`` is
still the repository's JSON-safe symbol mapping and ``merges`` is still
an ordered list of symbol-string pairs. The implementation is faster
because it:
* counts unique whitespace-normalized words once, weighted by frequency;
* stores a live inverted index from pair to word IDs;
* only rewrites words that contain the chosen best pair;
* uses a lazy-invalidated max heap instead of scanning every pair every
merge.
The merge loop is still sequential by nature, but initial word counting
can use multiple CPU workers while avoiding the repeated full-corpus
recount bottleneck.
"""
if num_merges < 0:
raise ValueError("num_merges must be non-negative")
if min_pair_frequency < 1:
raise ValueError("min_pair_frequency must be at least 1")
if chunk_size <= 0:
raise ValueError("chunk_size must be positive")
if isinstance(texts, str):
text_items = [texts]
else:
text_items = [str(text) for text in texts]
t0 = time.time()
word_freqs = self._count_word_frequencies_parallel(
text_items,
n_workers=n_workers,
chunk_size=chunk_size,
)
if verbose:
print(
f"[bpe] counted unique_words={len(word_freqs):,} "
f"total_words={sum(word_freqs.values()):,} in {time.time() - t0:.1f}s",
flush=True,
)
self._train_fast_from_word_frequencies(
word_freqs,
num_merges=num_merges,
min_pair_frequency=min_pair_frequency,
verbose=verbose,
start_time=t0,
)
def encode(self, text: str) -> List[int]:
"""Encode text into token IDs using the learned merge rules.
Byte-level tokenization guarantees that every UTF-8 byte has a base
token, so ordinary Unicode input can be encoded without an unknown
fallback. ``<unk>`` is still available for malformed model files or a
manually modified vocabulary.
Parameters
----------
text:
Raw Python string to encode.
Returns
-------
list[int]
Token IDs ready for model input.
"""
ids: List[int] = []
word_cache: Dict[str, List[int]] = {}
for word in self._preprocess_text(text):
cached = word_cache.get(word)
if cached is None:
symbols = self._word_to_symbols(word)
merged = self._apply_merges_to_word(symbols)
cached = [
self.vocab.get(symbol, self.special_tokens["<unk>"])
for symbol in merged
]
word_cache[word] = cached
ids.extend(cached)
return ids
def decode(self, ids: Sequence[int]) -> str:
"""Decode token IDs back into a Unicode string.
The decoder reconstructs bytes from byte and merged-byte symbols. Every
``</w>`` atom becomes a single space. The final output is stripped of a
trailing space introduced by the last word boundary.
Unknown IDs and non-boundary special tokens are skipped. If an ID maps
to bytes that are not valid UTF-8, Python's replacement character is
used rather than raising an exception.
"""
output = bytearray()
for token_id in ids:
symbol = self.id_to_symbol.get(int(token_id))
if symbol is None:
continue
if symbol in self.special_tokens and symbol != self.EOW_SYMBOL:
continue
for atom in self._symbol_to_atoms(symbol):
if atom == self.EOW_ATOM:
output.extend(b" ")
elif 0 <= atom <= 255:
output.append(atom)
return output.decode("utf-8", errors="replace").rstrip(" ")
def save(self, directory: str | Path) -> None:
"""Persist merge rules, vocabulary, and special tokens as JSON files.
Three files are written:
* ``vocab.json``: symbol string to ID mapping
* ``merges.json``: ordered list of BPE merge pairs
* ``special_tokens.json``: reserved special token IDs
The files are intentionally plain JSON so they can be reviewed, diffed,
archived, and audited without custom tooling.
"""
directory = Path(directory)
directory.mkdir(parents=True, exist_ok=True)
self._write_json(directory / "vocab.json", self.vocab)
self._write_json(directory / "merges.json", self.merges)
self._write_json(directory / "special_tokens.json", self.special_tokens)
@classmethod
def load(cls, directory: str | Path) -> "BPE_Tokenizer":
"""Restore a tokenizer previously written by :meth:`save`.
Parameters
----------
directory:
Directory containing ``vocab.json``, ``merges.json``, and
``special_tokens.json``.
Returns
-------
BPE_Tokenizer
Tokenizer with vocabulary, special tokens, merge rules, and merge
ranks restored.
"""
directory = Path(directory)
special_tokens = json.loads(
(directory / "special_tokens.json").read_text(encoding="utf-8")
)
tokenizer = cls(special_tokens={str(k): int(v) for k, v in special_tokens.items()})
vocab_data = json.loads((directory / "vocab.json").read_text(encoding="utf-8"))
tokenizer.vocab = {str(k): int(v) for k, v in vocab_data.items()}
tokenizer.id_to_symbol = {v: k for k, v in tokenizer.vocab.items()}
merges_data = json.loads((directory / "merges.json").read_text(encoding="utf-8"))
tokenizer.merges = [(str(left), str(right)) for left, right in merges_data]
tokenizer.merge_ranks = {
pair: rank for rank, pair in enumerate(tokenizer.merges)
}
return tokenizer
def vocabulary_size(self) -> int:
"""Return the number of known token symbols, including special tokens."""
return len(self.vocab)
@classmethod
def _validate_special_tokens(cls, special_tokens: MutableMapping[str, int]) -> None:
required = {"<unk>", "<pad>", "<s>", "</s>", cls.EOW_SYMBOL}
missing = required.difference(special_tokens)
if missing:
raise ValueError(f"Missing required special tokens: {sorted(missing)}")
if len(set(special_tokens.values())) != len(special_tokens):
raise ValueError("Special token IDs must be unique")
if special_tokens["<unk>"] != 0:
raise ValueError("This implementation reserves ID 0 for <unk>")
def _initialize_base_vocabulary(self) -> None:
"""Initialize special tokens and all 256 byte symbols."""
self.vocab.clear()
self.id_to_symbol.clear()
for token, token_id in sorted(self.special_tokens.items(), key=lambda item: item[1]):
self.vocab[token] = int(token_id)
self.id_to_symbol[int(token_id)] = token
next_id = max(self.special_tokens.values()) + 1
for byte_value in range(256):
self._add_symbol(self._atom_to_symbol(byte_value), preferred_id=next_id)
next_id += 1
def _add_symbol(self, symbol: Symbol, preferred_id: int | None = None) -> int:
"""Add a symbol to the vocabulary and return its ID."""
if symbol in self.vocab:
return self.vocab[symbol]
token_id = preferred_id if preferred_id is not None else self._next_available_id()
if token_id in self.id_to_symbol:
raise ValueError(f"Token ID collision for ID {token_id}")
self.vocab[symbol] = token_id
self.id_to_symbol[token_id] = symbol
return token_id
def _next_available_id(self) -> int:
return max(self.id_to_symbol, default=-1) + 1
@staticmethod
def _preprocess_text(text: str) -> List[str]:
"""Normalize whitespace and split text into words.
Consecutive whitespace characters, including spaces, tabs, and newlines,
are replaced by a single space. Leading and trailing whitespace is
removed. The returned list does not include explicit ``</w>`` strings;
that marker is appended in ``_word_to_symbols``.
"""
normalized = re.sub(r"\s+", " ", text.strip())
return [] if not normalized else normalized.split(" ")
def _build_training_corpus(self, text: str) -> Counter[Word]:
"""Create a frequency dictionary of tokenized words for BPE training."""
corpus: Counter[Word] = Counter()
for word in self._preprocess_text(text):
corpus[tuple(self._word_to_symbols(word))] += 1
return corpus
@staticmethod
def _count_word_frequencies_parallel(
texts: Sequence[str],
n_workers: int | None = None,
chunk_size: int = 512,
) -> Counter[str]:
"""Count normalized words, optionally using multiple worker processes."""
chunks = [texts[index : index + chunk_size] for index in range(0, len(texts), chunk_size)]
if not chunks:
return Counter()
workers = n_workers if n_workers is not None else max(1, (mp.cpu_count() or 2) - 1)
main_file = getattr(sys.modules.get("__main__"), "__file__", "")
if not main_file or main_file == "<stdin>":
workers = 1
if workers <= 1 or len(chunks) == 1:
total: Counter[str] = Counter()
for chunk in chunks:
total.update(_count_words_chunk(chunk))
return total
total = Counter()
with mp.Pool(processes=workers) as pool:
for counts in pool.imap_unordered(_count_words_chunk, chunks):
total.update(counts)
return total
def _train_fast_from_word_frequencies(
self,
word_freqs: Counter[str],
num_merges: int,
min_pair_frequency: int,
verbose: bool,
start_time: float,
) -> None:
"""Incrementally train BPE merges from weighted unique words."""
token_atoms: Dict[int, Tuple[int, ...]] = {
atom: (atom,) for atom in range(self.EOW_ATOM + 1)
}
next_internal_id = self.EOW_ATOM + 1
words: List[AtomWord] = []
freqs: List[int] = []
for word, frequency in word_freqs.items():
if frequency <= 0:
continue
words.append(tuple(word.encode("utf-8")) + (self.EOW_ATOM,))
freqs.append(int(frequency))
pair_counts: Dict[Tuple[int, int], int] = {}
pair_to_words: Dict[Tuple[int, int], set[int]] = {}
heap: List[Tuple[int, Tuple[int, int]]] = []
def bump(pair: Tuple[int, int], delta: int, word_id: int) -> None:
pair_counts[pair] = pair_counts.get(pair, 0) + delta
if pair_counts[pair] > 0:
pair_to_words.setdefault(pair, set()).add(word_id)
for word_id, atoms in enumerate(words):
frequency = freqs[word_id]
for pair, occurrences in self._count_atom_pairs(atoms).items():
bump(pair, occurrences * frequency, word_id)
for pair, count in pair_counts.items():
if count > 0:
heapq.heappush(heap, (-count, pair))
if verbose:
print(
f"[bpe] init unique_words={len(words):,} pairs={len(pair_counts):,} "
f"in {time.time() - start_time:.1f}s",
flush=True,
)
for step in range(num_merges):
best_pair: Tuple[int, int] | None = None
best_count = 0
while heap:
negative_count, pair = heapq.heappop(heap)
current_count = pair_counts.get(pair, 0)
if current_count == -negative_count and current_count > 0:
best_pair = pair
best_count = current_count
break
if best_pair is None or best_count < min_pair_frequency:
break
left, right = best_pair
new_internal_id = next_internal_id
next_internal_id += 1
token_atoms[new_internal_id] = token_atoms[left] + token_atoms[right]
left_symbol = self._atoms_to_symbol(token_atoms[left])
right_symbol = self._atoms_to_symbol(token_atoms[right])
new_symbol = self._atoms_to_symbol(token_atoms[new_internal_id])
if new_symbol not in self.vocab:
self._add_symbol(new_symbol)
merge_pair = (left_symbol, right_symbol)
self.merges.append(merge_pair)
self.merge_ranks[merge_pair] = len(self.merges) - 1
touched_pairs = set()
for word_id in list(pair_to_words.get(best_pair, ())):
frequency = freqs[word_id]
old_atoms = words[word_id]
old_pairs = self._count_atom_pairs(old_atoms)
if best_pair not in old_pairs:
continue
for pair, occurrences in old_pairs.items():
pair_counts[pair] = pair_counts.get(pair, 0) - occurrences * frequency
pair_words = pair_to_words.get(pair)
if pair_words is not None:
pair_words.discard(word_id)
touched_pairs.add(pair)
new_atoms = self._replace_atom_pair(old_atoms, left, right, new_internal_id)
words[word_id] = new_atoms
for pair, occurrences in self._count_atom_pairs(new_atoms).items():
bump(pair, occurrences * frequency, word_id)
touched_pairs.add(pair)
pair_counts.pop(best_pair, None)
pair_to_words.pop(best_pair, None)
for pair in touched_pairs:
count = pair_counts.get(pair, 0)
if count > 0:
heapq.heappush(heap, (-count, pair))
if verbose and ((step + 1) == 1 or (step + 1) % 500 == 0):
print(
f"[bpe] merge={step + 1:,}/{num_merges:,} "
f"best_count={best_count:,} vocab={self.vocabulary_size():,} "
f"elapsed={time.time() - start_time:.1f}s",
flush=True,
)
@staticmethod
def _count_atom_pairs(atoms: AtomWord) -> Counter[Tuple[int, int]]:
"""Count adjacent atom-token pairs in one encoded word."""
return Counter(zip(atoms[:-1], atoms[1:]))
@staticmethod
def _replace_atom_pair(
atoms: AtomWord,
left: int,
right: int,
new_id: int,
) -> AtomWord:
"""Replace non-overlapping ``(left, right)`` pairs in one atom word."""
output: List[int] = []
index = 0
while index < len(atoms):
if index < len(atoms) - 1 and atoms[index] == left and atoms[index + 1] == right:
output.append(new_id)
index += 2
else:
output.append(atoms[index])
index += 1
return tuple(output)
@classmethod
def _atoms_to_symbol(cls, atoms: Sequence[int]) -> Symbol:
"""Convert an atomic byte/EOW sequence to the tokenizer symbol format."""
if tuple(atoms) == (cls.EOW_ATOM,):
return cls.EOW_SYMBOL
return " ".join(str(atom) for atom in atoms)
def _word_to_symbols(self, word: str) -> List[Symbol]:
"""Convert one word into byte symbols followed by ``</w>``."""
symbols = [self._atom_to_symbol(byte_value) for byte_value in word.encode("utf-8")]
symbols.append(self.EOW_SYMBOL)
return symbols
@classmethod
def _atom_to_symbol(cls, atom: int) -> Symbol:
"""Convert an atomic byte or end-of-word atom to a symbol string."""
if atom == cls.EOW_ATOM:
return cls.EOW_SYMBOL
if not 0 <= atom <= 255:
raise ValueError(f"Invalid atom: {atom}")
return str(atom)
@classmethod
def _symbol_to_atoms(cls, symbol: Symbol) -> Tuple[int, ...]:
"""Convert a symbol string back to its atomic byte/EOW sequence."""
if symbol == cls.EOW_SYMBOL:
return (cls.EOW_ATOM,)
try:
return tuple(int(part) for part in symbol.split(" "))
except ValueError as exc:
raise ValueError(f"Invalid symbol encoding: {symbol!r}") from exc
@classmethod
def _merge_symbol(cls, pair: Pair) -> Symbol:
"""Return the canonical symbol produced by merging a pair."""
atoms = cls._symbol_to_atoms(pair[0]) + cls._symbol_to_atoms(pair[1])
if atoms == (cls.EOW_ATOM,):
return cls.EOW_SYMBOL
return " ".join(str(atom) for atom in atoms)
@staticmethod
def _count_pairs(corpus: Counter[Word]) -> Counter[Pair]:
"""Count adjacent symbol pairs, weighted by word frequency."""
pair_counts: Counter[Pair] = Counter()
for symbols, frequency in corpus.items():
if len(symbols) < 2:
continue
for index in range(len(symbols) - 1):
pair_counts[(symbols[index], symbols[index + 1])] += frequency
return pair_counts
@staticmethod
def _best_pair(pair_counts: Counter[Pair]) -> Tuple[Pair, int]:
"""Select the highest-frequency pair with deterministic tie-breaking."""
best_pair, best_count = max(
pair_counts.items(),
key=lambda item: (item[1], item[0][0], item[0][1]),
)
return best_pair, best_count
@staticmethod
def _replace_pair_in_corpus(
corpus: Counter[Word],
pair: Pair,
new_symbol: Symbol,
) -> Counter[Word]:
"""Replace all non-overlapping occurrences of a pair in the corpus."""
updated: Counter[Word] = Counter()
left, right = pair
for symbols, frequency in corpus.items():
merged: List[Symbol] = []
index = 0
while index < len(symbols):
if (
index < len(symbols) - 1
and symbols[index] == left
and symbols[index + 1] == right
):
merged.append(new_symbol)
index += 2
else:
merged.append(symbols[index])
index += 1
updated[tuple(merged)] += frequency
return updated
def _apply_merges_to_word(self, symbols: List[Symbol]) -> List[Symbol]:
"""Apply learned merge rules to one word.
The naive approach scans every learned merge rule for every word. That
is easy to understand but slow for corpus-scale encoding. This method
instead finds the currently present adjacent pair with the best
training rank, merges it, and repeats until no learned pair remains.
The result is equivalent to applying merge rules in training order, but
it usually touches only the pairs that can actually occur in the word.
"""
current = list(symbols)
while len(current) >= 2:
best_pair: Pair | None = None
best_rank: int | None = None
for index in range(len(current) - 1):
pair = (current[index], current[index + 1])
rank = self.merge_ranks.get(pair)
if rank is not None and (best_rank is None or rank < best_rank):
best_pair = pair
best_rank = rank
if best_pair is None:
break
current = self._replace_pair_in_word(
current,
best_pair,
self._merge_symbol(best_pair),
)
return current
@staticmethod
def _replace_pair_in_word(
symbols: Sequence[Symbol],
pair: Pair,
new_symbol: Symbol,
) -> List[Symbol]:
"""Replace all non-overlapping occurrences of a pair in one word."""
merged: List[Symbol] = []
left, right = pair
index = 0
while index < len(symbols):
if (
index < len(symbols) - 1
and symbols[index] == left
and symbols[index + 1] == right
):
merged.append(new_symbol)
index += 2
else:
merged.append(symbols[index])
index += 1
return merged
@staticmethod
def _write_json(path: Path, data: object) -> None:
"""Write stable, human-readable JSON."""
path.write_text(
json.dumps(data, ensure_ascii=False, indent=2, sort_keys=True) + "\n",
encoding="utf-8",
)
def _demo() -> None:
"""Train a tiny Leipzig-focused tokenizer and show a round trip."""
sample_corpus = """
Leipzig ist eine Stadt mit Musik, Messe und Auenwald.
Die Thomaskirche und der Thomanerchor prägen Leipzig.
Im Leipziger Auenwald treffen Spaziergänge, Wasser und Geschichte zusammen.
Bach wirkte in Leipzig, und viele Gäste besuchen die Thomaskirche.
"""
tokenizer = BPE_Tokenizer()
tokenizer.train(sample_corpus, num_merges=80)
example = "Leipzig und die Thomaskirche liegen nahe am Auenwald."
encoded = tokenizer.encode(example)
decoded = tokenizer.decode(encoded)
print("Vocabulary size:", tokenizer.vocabulary_size())
print("Merge rules learned:", len(tokenizer.merges))
print("Input: ", example)
print("Encoded:", encoded)
print("Decoded:", decoded)
print("Round trip OK:", decoded == example)
with tempfile.TemporaryDirectory(prefix="bpe_tokenizer_demo_") as tmp_dir:
tokenizer.save(tmp_dir)
restored = BPE_Tokenizer.load(tmp_dir)
restored_decoded = restored.decode(restored.encode(example))
print("Saved files:", tmp_dir)
print("Reload round trip OK:", restored_decoded == example)
if __name__ == "__main__":
_demo()
|