外观
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")