Tokenwerk · The LLM textbook

Chapter 33 · VII · Build your LLM · 10 minutes

Step 1: The tokenizer

The first component of your model: a class that splits text into token IDs and puts it back together, with real byte-level BPE like GPT-2.

What we build

A language model only computes with numbers. So the first component translates text into token IDs and back again: the tokenizer. You learned the idea in chapter 12. Here it is in three sentences:

  • We start with the 256 possible bytes. They can represent any text in the world.
  • During training we look for the most frequent neighboring pair and merge it into a new token. We repeat that until the vocabulary is big enough; for us, 1024 tokens.
  • During encoding we apply these merge rules to new text in the same order.

At the end of this step you have the file tokenizer.py with a class Tokenizer that can do all of this.

A trick up front: split into words first

If we counted pairs across the whole text at once, it would be pretty slow with millions of tokens and hundreds of merge rounds. So we first split the text into words, punctuation and whitespace, and then count each word only once, with its frequency. The word “ and” occurs thousands of times but only needs to be processed once.

Side effect: merges only happen inside a word. So a token can never be half “cat” and half the next word. GPT-2 and many other models do exactly the same.

python
# Byte-level BPE tokenizer: text -> token IDs -> text (chapter 12)
import json
import re
from collections import Counter

# Split text into words, punctuation and whitespace first. Merges only happen inside these pieces.
PATTERN = re.compile(r" ?\w+| ?[^\w\s]+|\s+")
SPECIALS = ["<pad>", "<bos>", "<eos>", "<user>", "<assistant>"]

The regular expression PATTERN describes how to split: an optional space plus letters ( ?\w+), an optional space plus punctuation ( ?[^\w\s]+) or whitespace (\s+). “Da sprach er:” becomes ["Da", " sprach", " er", ":"], and “Lily said:” becomes ["Lily", " said", ":"]. The space stays with the following word, which is why token lists often contain tokens like “ said” with a leading space.

You know the five special tokens from chapter 12: <bos> and <eos> mark the start and end, <pad> fills up, and we need <user> and <assistant> later for the chat.

Merging a pair

The most important helper: it replaces every occurrence of a pair in a list of IDs with a new ID. From left to right and without overlaps: with the pair (a, a), [a, a, a] becomes [aa, a].

python
def merge(ids, pair, new_id):
    """Replace every occurrence of pair in ids with new_id (left to right, no overlaps)."""
    out, i = [], 0
    while i < len(ids):
        if i + 1 < len(ids) and (ids[i], ids[i + 1]) == pair:
            out.append(new_id)
            i += 2
        else:
            out.append(ids[i])
            i += 1
    return out

The class: building the vocabulary

A tokenizer really consists only of its list of merge rules. Everything else can be computed from it:

python
class Tokenizer:
    def __init__(self, merges=()):
        self.merges = [tuple(pair) for pair in merges]
        self.pieces = [bytes([b]) for b in range(256)]           # IDs 0-255: the bytes themselves
        self.rank = {}
        for pair in self.merges:                                   # every merge rule gets the next ID
            self.rank[pair] = len(self.pieces)
            self.pieces.append(self.pieces[pair[0]] + self.pieces[pair[1]])
        self.special = {name: len(self.pieces) + i for i, name in enumerate(SPECIALS)}
        self.cache = {}

    @property
  • pieces is the vocabulary as bytes: IDs 0 to 255 are the individual bytes, followed by the combined piece for each merge rule. If rule 1 is “t” + “he”, the new piece is simply b"t" + b"he".
  • rank remembers which ID each pair received. The order matters: whatever was learned earlier is also applied first later.
  • We append the special tokens at the end, so they don't get in the way of the byte IDs.
  • @property turns vocab_size into a property: you write tok.vocab_size without parentheses.

Training: finding the most frequent pair

python
class Tokenizer:
    @classmethod

Step by step:

  1. Counter(PATTERN.findall(text)) splits the text and counts how often each word occurs.
  2. seqs holds the current split of every word, byte by byte at the start.
  3. In every round we count all neighboring pairs, weighted by word frequency.
  4. The most frequent pair wins. On a tie, the smaller ID decides, so the result is always the same.
  5. If the best pair occurs only once, we stop: it isn't worth its own token.
  6. Otherwise we remember the rule and apply it to all words.

The condition in while makes sure exactly vocab_size tokens come out in the end: 256 bytes plus the merges plus the 5 special tokens.

Encoding and decoding

python
class Tokenizer:
    def encode_word(self, word):
        if word in self.cache:
            return self.cache[word]
        ids = list(word.encode("utf-8"))
        while len(ids) > 1:
            # the rule with the lowest rank first – exactly the order from training
            pair = min(zip(ids, ids[1:]), key=lambda p: self.rank.get(p, float("inf")))
            if pair not in self.rank:
                break
            ids = merge(ids, pair, self.rank[pair])
        self.cache[word] = ids
        return ids

    def encode(self, text):
        ids = []
        for word in PATTERN.findall(text):
            ids.extend(self.encode_word(word))
        return ids

    def decode(self, ids):
        data = b"".join(self.pieces[i] for i in ids if i < len(self.pieces))   # skip special tokens
        return data.decode("utf-8", errors="replace")

encode_word applies the learned rules to a word: among all pairs in the word it looks for the one with the lowest rank, i.e. the rule learned first in training, and merges it. This repeats until no rule fits anymore. Because the same words keep coming back, we store the result in cache. That makes encoding many times faster.

decode goes the other way: join the bytes of the tokens and read them as UTF-8. Special tokens are skipped. errors="replace" catches broken byte sequences in case a model ever produces an incomplete character (chapter 12).

Saving and loading

python
class Tokenizer:
    def save(self, path):
        with open(path, "w", encoding="utf-8") as f:
            json.dump({"merges": self.merges}, f)

    @classmethod

Only the merge rules are saved as JSON. Nothing more is needed; __init__ recomputes the rest when loading. @classmethod means the method belongs to the class, not to a single object. That is why you later write Tokenizer.load("…") instead of building an empty object first.

The whole file

python
# Byte-level BPE tokenizer: text -> token IDs -> text (chapter 12)
import json
import re
from collections import Counter

# Split text into words, punctuation and whitespace first. Merges only happen inside these pieces.
PATTERN = re.compile(r" ?\w+| ?[^\w\s]+|\s+")
SPECIALS = ["<pad>", "<bos>", "<eos>", "<user>", "<assistant>"]


def merge(ids, pair, new_id):
    """Replace every occurrence of pair in ids with new_id (left to right, no overlaps)."""
    out, i = [], 0
    while i < len(ids):
        if i + 1 < len(ids) and (ids[i], ids[i + 1]) == pair:
            out.append(new_id)
            i += 2
        else:
            out.append(ids[i])
            i += 1
    return out


class Tokenizer:
    def __init__(self, merges=()):
        self.merges = [tuple(pair) for pair in merges]
        self.pieces = [bytes([b]) for b in range(256)]           # IDs 0-255: the bytes themselves
        self.rank = {}
        for pair in self.merges:                                   # every merge rule gets the next ID
            self.rank[pair] = len(self.pieces)
            self.pieces.append(self.pieces[pair[0]] + self.pieces[pair[1]])
        self.special = {name: len(self.pieces) + i for i, name in enumerate(SPECIALS)}
        self.cache = {}

    @property
    def vocab_size(self):
        return len(self.pieces) + len(SPECIALS)

    @classmethod
    def train(cls, text, vocab_size):
        words = Counter(PATTERN.findall(text))                   # every word only once, with its count
        seqs = {w: list(w.encode("utf-8")) for w in words}
        merges = []
        while 256 + len(merges) + len(SPECIALS) < vocab_size:
            pairs = Counter()
            for w, count in words.items():
                s = seqs[w]
                for pair in zip(s, s[1:]):
                    pairs[pair] += count
            if not pairs:
                break
            pair, count = max(pairs.items(), key=lambda item: (item[1], -item[0][0], -item[0][1]))
            if count < 2:
                break
            new_id = 256 + len(merges)
            merges.append(pair)
            for w in words:
                seqs[w] = merge(seqs[w], pair, new_id)
        return cls(merges)

    def encode_word(self, word):
        if word in self.cache:
            return self.cache[word]
        ids = list(word.encode("utf-8"))
        while len(ids) > 1:
            # the rule with the lowest rank first – exactly the order from training
            pair = min(zip(ids, ids[1:]), key=lambda p: self.rank.get(p, float("inf")))
            if pair not in self.rank:
                break
            ids = merge(ids, pair, self.rank[pair])
        self.cache[word] = ids
        return ids

    def encode(self, text):
        ids = []
        for word in PATTERN.findall(text):
            ids.extend(self.encode_word(word))
        return ids

    def decode(self, ids):
        data = b"".join(self.pieces[i] for i in ids if i < len(self.pieces))   # skip special tokens
        return data.decode("utf-8", errors="replace")

    def save(self, path):
        with open(path, "w", encoding="utf-8") as f:
            json.dump({"merges": self.merges}, f)

    @classmethod
    def load(cls, path):
        with open(path, encoding="utf-8") as f:
            return cls(json.load(f)["merges"])

Try it out

Start a Python console in the my-llm folder with python and try the tokenizer on a tiny text:

python
from tokenizer import Tokenizer
tok = Tokenizer.train("the cat and the can and the cat", 270)
print(tok.vocab_size)
ids = tok.encode("the cat")
print(ids)
print([tok.decode([i]) for i in ids])
print(tok.decode(ids) == "the cat")

Output:

text
270
[259, 263]
['the', ' cat']
True

270 tokens means: 256 bytes, 5 special tokens, which leaves 9 merges. They were enough to turn “the” and “ cat” into whole tokens. Try 265 instead of 270: then only 4 merges remain, and “ cat” falls apart into “ c”, “a”, “t”. And most importantly: the way back gives exactly the original text.

If something goes wrong

  • ModuleNotFoundError: No module named 'tokenizer': the Python console isn't running in the my-llm folder. Switch into it with cd my-llm.
  • IndentationError: the indentation slipped while typing. Python is strict about this: methods of a class are indented by exactly 4 spaces.
  • Different IDs than in the book: no problem, as long as True is at the end. The exact numbers depend on the training text.

Key points

  • The tokenizer learns merge rules: the most frequent pair becomes a new token, again and again.
  • Splitting into words first makes training fast and prevents tokens across word boundaries.
  • Only the rules are saved. Encoding applies them in the same order; decoding joins bytes together.