bpe.py
Byte pair encoding — the reference for /vocabulary/
119 lines. This is the file the build actually runs, copied verbatim at build time.
- 1
"""Byte pair encoding, trained from scratch — the reference implementation. - 2
- 3
This is the teaching version of the algorithm that produced the vocabularies on - 4
/tokens/. It mirrors trainer.js exactly, symbol for symbol, and checkbpe.py - 5
fails the build if the two ever diverge. - 6
- 7
Two honest simplifications versus a production tokenizer: - 8
- 9
* It merges characters, not UTF-8 bytes. Real tokenizers start from the 256 - 10
possible bytes so that any input at all is representable. Starting from - 11
characters is easier to watch and makes no difference to the mechanism. - 12
* Pre-tokenization is one simple rule (optional whitespace, then a run of - 13
non-whitespace) rather than a long regex. That rule is why leading spaces - 14
end up welded to words — which is the point worth seeing. - 15
- 16
Tie-breaking: highest pair count wins; ties go to whichever pair was seen first - 17
in a left-to-right scan. First-seen order avoids comparing strings, which - 18
Python and JavaScript order differently for astral characters. - 19
""" - 20
- 21
import re - 22
- 23
PRETOKEN = re.compile(r"\s*\S+|\s+") - 24
- 25
- 26
def pretokenize(text): - 27
"""Split into chunks of optional leading whitespace plus a run of non-space.""" - 28
return PRETOKEN.findall(text) - 29
- 30
- 31
def word_counts(text): - 32
"""Unique chunks and how often each occurs, in first-seen order.""" - 33
counts = {} - 34
for chunk in pretokenize(text): - 35
counts[chunk] = counts.get(chunk, 0) + 1 - 36
return counts - 37
- 38
- 39
def _pair_stats(words): - 40
"""pair -> total count, in first-seen order (dicts preserve insertion order, - 41
and so do JavaScript Maps — that is what makes tie-breaking identical).""" - 42
counts = {} - 43
for symbols, freq in words: - 44
for i in range(len(symbols) - 1): - 45
pair = (symbols[i], symbols[i + 1]) - 46
counts[pair] = counts.get(pair, 0) + freq - 47
return counts - 48
- 49
- 50
def _merge_in(symbols, a, b): - 51
"""Replace every adjacent a,b with the joined symbol.""" - 52
out = [] - 53
i = 0 - 54
n = len(symbols) - 55
while i < n: - 56
if i < n - 1 and symbols[i] == a and symbols[i + 1] == b: - 57
out.append(a + b) - 58
i += 2 - 59
else: - 60
out.append(symbols[i]) - 61
i += 1 - 62
return out - 63
- 64
- 65
def train(text, num_merges): - 66
"""Run BPE. Returns the merge list and a step-by-step record.""" - 67
words = [(list(chunk), freq) for chunk, freq in word_counts(text).items()] - 68
- 69
alphabet = {} - 70
for chunk in word_counts(text): - 71
for ch in chunk: - 72
alphabet[ch] = True - 73
- 74
merges, steps = [], [] - 75
for _ in range(num_merges): - 76
counts = _pair_stats(words) - 77
if not counts: - 78
break - 79
# Strict >, scanning in first-seen order: the earliest of any tied pairs - 80
# wins. No string comparison, so Python and JavaScript cannot disagree. - 81
best = None - 82
for pair, count in counts.items(): - 83
if best is None or count > counts[best]: - 84
best = pair - 85
if counts[best] < 2: - 86
break # nothing repeats; merging is pointless - 87
- 88
a, b = best - 89
words = [(_merge_in(sym, a, b), freq) for sym, freq in words] - 90
merges.append((a, b)) - 91
steps.append({ - 92
"pair": [a, b], - 93
"token": a + b, - 94
"count": counts[best], - 95
"vocab": len(alphabet) + len(merges), - 96
}) - 97
- 98
return { - 99
"merges": merges, - 100
"steps": steps, - 101
"alphabet": sorted(alphabet), - 102
"vocab_size": len(alphabet) + len(merges), - 103
} - 104
- 105
- 106
def encode(word, merges): - 107
"""Apply a trained merge list to one chunk, in the order learned.""" - 108
symbols = list(word) - 109
for a, b in merges: - 110
symbols = _merge_in(symbols, a, b) - 111
return symbols - 112
- 113
- 114
def encode_text(text, merges): - 115
out = [] - 116
for chunk in pretokenize(text): - 117
out.extend(encode(chunk, merges)) - 118
return out