precompute_merges.py
Computes the figures on /vocabulary/
66 lines. This is the file the build actually runs, copied verbatim at build time.
- 1
"""Generate merges/data.json — every figure on /merges/, computed not written.""" - 2
- 3
import json - 4
import os - 5
- 6
import bpe - 7
from cjsload import Reference - 8
from corpora import BERRIES - 9
- 10
OUT = os.path.join(os.path.dirname(os.path.abspath(__file__)), "..", - 11
"vocabulary", "data.json") - 12
- 13
N_MERGES = 30 - 14
- 15
# Words that show the three outcomes: learned whole, learned in pieces, unseen. - 16
PROBES = [" strawberry", "strawberry", " blackberry", "strawberries", " kiwi"] - 17
- 18
# Where the assembly of " strawberry" visibly changes. - 19
CHECKPOINTS = [0, 1, 2, 3, 4, 16, 17] - 20
- 21
- 22
def main(): - 23
trained = bpe.train(BERRIES, N_MERGES) - 24
merges = trained["merges"] - 25
real = Reference("o200k_base") - 26
- 27
data = { - 28
"corpus": BERRIES, - 29
"corpus_chars": len(BERRIES), - 30
"corpus_words": len(bpe.pretokenize(BERRIES)), - 31
"corpus_unique": len(set(bpe.pretokenize(BERRIES))), - 32
"alphabet": trained["alphabet"], - 33
"steps": trained["steps"], - 34
"vocab_size": trained["vocab_size"], - 35
"n_merges": len(merges), - 36
"checkpoints": [ - 37
{"after": n, "symbols": bpe.encode(" strawberry", merges[:n])} - 38
for n in CHECKPOINTS - 39
], - 40
"probes": [ - 41
{ - 42
"text": p, - 43
"toy": bpe.encode(p, merges), - 44
"real": [t["text"] for t in real.tokens(p)], - 45
} - 46
for p in PROBES - 47
], - 48
} - 49
- 50
with open(OUT, "w", encoding="utf-8") as fh: - 51
json.dump(data, fh, ensure_ascii=False, indent=1) - 52
- 53
print(f"wrote {os.path.relpath(OUT, os.path.dirname(OUT))}") - 54
print(f" corpus: {data['corpus_chars']} chars, {data['corpus_words']} words, " - 55
f"{data['corpus_unique']} unique") - 56
print(f" {data['n_merges']} merges, alphabet {len(data['alphabet'])}, " - 57
f"vocab {data['vocab_size']}") - 58
print(" first merges: " + - 59
" ".join(f"{s['token']!r}" for s in data["steps"][:5])) - 60
for p in data["probes"]: - 61
print(f" {p['text']!r:>16} toy {len(p['toy']):>2} real {len(p['real']):>2}") - 62
- 63
- 64
if __name__ == "__main__": - 65
main()