从 Karpathy 手写 BPE 到 CS336 `train_bpe` 作业实现

这次 train_bpe 的实现过程,大体分成两个阶段:

  1. 先跟着 Karpathy 的 tokenizer/BPE 手写思路,实现最核心的 BPE 训练逻辑。
  2. 再根据作业测试要求,对接口、特殊 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 转发到自己的实现。

同时,作业还要求先进行文本预处理:

  1. 根据 special_tokens 分割文本。
  2. 再对普通文本使用 GPT-2 的正则进行预分词。
  3. 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,最后再统一把 vocabmerges 转成 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 和速度的要求。

相关推荐
无忧智库1 小时前
产业数字化平台建设方案:从单点工厂改造到城市产业生态的全景落地指南(PPT)
大数据·人工智能
wuhanzhanhui1 小时前
轻盈的力量:2026 武汉国际发泡材料技术工业展览会,重构未来工业的“呼吸感”
大数据·人工智能
Ai_easygo1 小时前
多Agent数据分析报告自动化实战:用CrewAI组支AI团队,丢份数据就出报告(从架构到评估)
人工智能·数据分析·自动化
集芯微电科技有限公司1 小时前
替代LM4890音频功率放大器低电磁干扰辐射
人工智能·单片机·嵌入式硬件·生成对抗网络·计算机外设
开开心心就好1 小时前
Word双击预览图片插件弥补Word功能缺失
人工智能·python·智能手机·ocr·电脑·word·音视频
战场小包1 小时前
世界杯结束了,我用 AI 造了平行宇宙,这次结局你写
前端·人工智能·ai编程
冻感糕人~1 小时前
大模型学习指南:收藏这份AI Agent四层工程地图(小白程序员必备)
java·大数据·人工智能·学习·大模型·agent·大模型学习
问商十三载1 小时前
2026工业制造AI引擎生成式优化怎么做?3层适配法避通用坑,零成本提31%抓取权重附校验清单
大数据·人工智能·制造
手写码匠2 小时前
Android 17 灵魂拷问深度解析:隐私、大屏、AI 端侧全面适配实战
人工智能·深度学习·算法·aigc