Skip to content

llm/bpe.py ​

配套源码,运行方法见同目录 README。返回实验总览。

python
"""Readable byte BPE. Original teaching code, not CS336's official solution.

Boundaries: split into runs of whitespace/non-whitespace. No merge crosses a
run or a document. The tie rule is lexicographically smallest integer pair.
This intentionally differs from the official GPT-2 regex and tie rule.
"""
from collections import Counter
import json
import re


def merge_pair(ids, pair, new_id):
    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 ByteBPE:
    def __init__(self, merges=()):
        self.merges = [tuple(pair) for pair in merges]
        self.vocab = {i: bytes([i]) for i in range(256)}
        for token_id, (left, right) in enumerate(self.merges, 256):
            self.vocab[token_id] = self.vocab[left] + self.vocab[right]

    @staticmethod
    def pieces(text):
        return [list(s.encode("utf-8")) for s in re.findall(r"\s+|\S+", text)]

    @classmethod
    def train(cls, documents, num_merges=32):
        if num_merges < 0:
            raise ValueError("num_merges must be nonnegative")
        pieces = [p for doc in documents for p in cls.pieces(doc)]
        merges = []
        for token_id in range(256, 256 + num_merges):
            counts = Counter(pair for p in pieces for pair in zip(p, p[1:]))
            if not counts:
                break
            pair = min(counts, key=lambda pair: (-counts[pair], pair))
            merges.append(pair)
            pieces = [merge_pair(p, pair, token_id) for p in pieces]
        return cls(merges)

    def encode(self, text):
        result = []
        for piece in self.pieces(text):
            for new_id, pair in enumerate(self.merges, 256):
                piece = merge_pair(piece, pair, new_id)
            result.extend(piece)
        return result

    def decode(self, ids):
        # An arbitrary generated prefix may end inside a UTF-8 code point.
        return b"".join(self.vocab[int(i)] for i in ids).decode("utf-8", errors="replace")

    def to_dict(self):
        return {"kind": "bpe", "merges": self.merges}

    def save(self, path):
        path.write_text(json.dumps(self.to_dict()), encoding="utf-8")

    @classmethod
    def load(cls, path):
        return cls(json.loads(path.read_text(encoding="utf-8"))["merges"])


class CharacterTokenizer:
    """Small-vocabulary baseline; unknown characters deliberately raise errors."""
    def __init__(self, characters):
        self.characters = list(characters)
        self.vocab = dict(enumerate(self.characters))
        self.mapping = {c: i for i, c in self.vocab.items()}

    @classmethod
    def train(cls, documents):
        return cls(sorted(set("".join(documents))))

    def encode(self, text):
        return [self.mapping[c] for c in text]

    def decode(self, ids):
        return "".join(self.vocab[int(i)] for i in ids)

    def to_dict(self):
        return {"kind": "char", "characters": self.characters}


def tokenizer_from_dict(data):
    if data["kind"] == "bpe":
        return ByteBPE(data["merges"])
    if data["kind"] == "char":
        return CharacterTokenizer(data["characters"])
    raise ValueError("unknown tokenizer kind")