这次 train_bpe 的实现过程,大体分成两个阶段:
- 先跟着 Karpathy 的 tokenizer/BPE 手写思路,实现最核心的 BPE 训练逻辑。
- 再根据作业测试要求,对接口、特殊 token、GPT-2 预分词、输出格式和速度进行适配。
参考视频是 Karpathy 的手写 tokenizer 讲解:Bilibili 视频链接。
1. 跟着 Karpathy 实现 BPE 核心逻辑
最开始的核心代码来自 Karpathy 视频里的简化版 BPE 思路:
get_stats:统计相邻 token pair 的出现频率。merge:把最高频 pair 合并成新的 token id。- 训练循环:每一轮选择最高频 pair,生成一个新的 token id,并更新 token 序列。
这个版本能帮助理解 BPE 的本质:
BPE 并不是一开始就按词切分,而是从 byte 级别开始,不断把高频相邻 pair 合并成更大的 token。
最初的训练内部主要用 int 表示 token id,例如:
python
"the".encode("utf-8") -> [116, 104, 101]
然后每次 merge 会生成新的 id:
python
(116, 104) -> 257
这个思路适合训练过程,但还不能直接满足作业测试。
2. 根据作业要求做适配
CS336 作业不是只要求写一个 toy BPE,而是要求实现测试接口:
python
run_train_bpe(input_path, vocab_size, special_tokens)
并返回:
python
vocab, merges
其中:
python
vocab: dict[int, bytes]
merges: list[tuple[bytes, bytes]]
所以第一步是把 tests/adapters.py 里的 run_train_bpe 转发到自己的实现。
同时,作业还要求先进行文本预处理:
- 根据
special_tokens分割文本。 - 再对普通文本使用 GPT-2 的正则进行预分词。
- BPE 只能在每个 pre-token 内部合并,不能跨 pre-token 或跨 special token 边界合并。
因此代码里使用了 GPT-2 风格正则:
python
gpt2pat = re.compile(
r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""
)
然后每个 pre-token 单独编码成 byte 序列参与训练。
3. vocab 和 merges 的 bytes 输出问题
一开始的想法是:训练时全部用 int,最后再统一把 vocab 和 merges 转成 bytes。
但后来发现这样会影响 BPE 的 tie-break 规则。
测试要求:如果多个 pair 的频率相同,要按 pair 对应的 bytes 进行比较,而不是按 token id 比较。
所以最终选择在训练过程中同步维护:
python
vocab: int -> bytes
初始 vocab 是:
python
vocab = {i: bytes([i]) for i in range(256)}
special token 也加入 vocab:
python
vocab[256 + i] = st.encode("utf-8")
每次 merge 后立刻更新:
python
vocab[new_idx] = vocab[merge_pair[0]] + vocab[merge_pair[1]]
这样下一轮如果 pair 中包含新 token,也可以通过 vocab 找到它对应的真实 bytes。
选择最高频 pair 时使用:
python
merge_pair = max(stats, key=lambda p: (stats[p], vocab[p[0]], vocab[p[1]]))
这解决了 merges 和参考答案不一致的问题。
4. 速度优化
正确性通过后,速度测试仍然失败。朴素版本每一轮都会扫描完整 token 列表,速度太慢。
观察测试数据后发现,语料中有大量重复 pre-token。于是把普通列表改成 Counter:
python
unique_token_list = Counter(tuple(text.encode("utf-8")) for text in new_text_list)
这样保存的是:
python
token序列 -> 出现次数
统计 pair 时按出现次数加权:
python
stats[pair] = stats.get(pair, 0) + count
这一步大幅减少了重复扫描。
后来速度仍然接近 1.5 秒阈值,于是继续优化 merge:
每一轮真正包含目标 pair 的 token 很少,所以先检查 token 中是否存在目标 pair。如果不存在,就直接复用原 token,避免重新创建列表。
最后还把:
python
old_ids[0], old_ids[1]
提前解包成局部变量:
python
old_idx1, old_idx2 = old_ids
减少循环中的下标访问开销。这个优化很小,但在速度测试压线时有帮助。
最终测试结果:
text
test_train_bpe_speed PASSED
test_train_bpe PASSED
test_train_bpe_special_tokens PASSED
最终代码
tests/adapters.py 中的转发代码:
python
def run_train_bpe(
input_path,
vocab_size,
special_tokens,
**kwargs,
):
from cs336_basics.tokenizer import run_train_bpe as train_bpe
return train_bpe(
input_path=input_path,
vocab_size=vocab_size,
special_tokens=special_tokens,
)
核心实现代码:
python
import regex as re
from collections import Counter
def get_stats(unique_token_list):
stats = {}
for token, count in unique_token_list.items():
for i in range(len(token) - 1):
pair = (token[i], token[i + 1])
stats[pair] = stats.get(pair, 0) + count
return stats
def merge(old_ids, unique_token_list, new_idx):
new_token_list = Counter()
old_idx1, old_idx2 = old_ids
for token, count in unique_token_list.items():
found = False
for i in range(len(token) - 1):
if token[i] == old_idx1 and token[i + 1] == old_idx2:
found = True
break
if not found:
new_token_list[token] += count
continue
new_token = []
i = 0
while i < len(token):
if i < len(token) - 1 and token[i] == old_idx1 and token[i + 1] == old_idx2:
new_token.append(new_idx)
i += 2
else:
new_token.append(token[i])
i += 1
new_token_list[tuple(new_token)] += count
return new_token_list
def run_train_bpe(input_path, vocab_size, special_tokens):
with open(input_path, "r", encoding="utf-8") as f:
text = f.read()
escaped_special_tokens = []
for st in special_tokens:
escaped_special_tokens.append(re.escape(st))
split_pattern = "|".join(escaped_special_tokens)
text_list = re.split(split_pattern, text)
gpt2pat = re.compile(
r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+"""
)
new_text_list = []
for text in text_list:
new_text_list.extend(re.findall(gpt2pat, text))
unique_token_list = Counter(tuple(text.encode("utf-8")) for text in new_text_list)
vocab = {i: bytes([i]) for i in range(256)}
for i, st in enumerate(special_tokens):
vocab[256 + i] = st.encode("utf-8")
num_merges = vocab_size - 256 - len(special_tokens)
merges = []
for i in range(num_merges):
stats = get_stats(unique_token_list)
if not stats:
break
merge_pair = max(stats, key=lambda p: (stats[p], vocab[p[0]], vocab[p[1]]))
new_idx = 256 + len(special_tokens) + i
vocab[new_idx] = vocab[merge_pair[0]] + vocab[merge_pair[1]]
merges.append((vocab[merge_pair[0]], vocab[merge_pair[1]]))
unique_token_list = merge(merge_pair, unique_token_list, new_idx)
return vocab, merges
总结
这次作业的核心收获是:Karpathy 视频里的 BPE 代码适合理解算法本质,但要通过课程测试,还需要做工程化适配。
主要修改包括:
- 按作业接口接入
adapters.py。 - 根据 special token 和 GPT-2 正则进行预分词。
- 保证不同 pre-token 之间不会发生 merge。
- 内部训练继续使用 int token id。
vocab从一开始维护成int -> bytes。- 用 bytes 元组完成频率相同时的 tie-break。
- 用
Counter去重并按出现次数加权统计。 - merge 时跳过不包含目标 pair 的 token,提高速度。
最终实现既保持了 Karpathy 手写 BPE 的核心结构,又满足了作业对正确性、special token 和速度的要求。