数据截至 (上游 commit d5827816baed)
02 · BPE 模型与训练器
这一章讲什么: BPE 的两个半算法——推理(给定 merges 表,把一个词切成子词)与训练(从语料学出 vocab + merges)。教学版 BPE 二十行就能写完,但要在 TB 级语料上训、在每毫秒级延迟下跑,需要本章这些数据结构:链表、优先队列、缓存、增量计数。读完你会知道
Word::merge_all为什么是 tokenizers 全库最精巧的一段代码。
1. 它要解决的小问题
BPE 推理的规则一句话:把词先切成字符,然后反复合并「在 merges 表里排名最靠前」的相邻对,直到没有可合并的对。
朴素实现每轮都扫一遍全词找最高优先级对——词长 n、合并 n 轮,就是 O(n²),而且每轮都要重建字符串。GPT-2 时代的 Python 实现就是这么干的,长词(URL、中文长串)直接拖死吞吐。
训练侧反过来: 从语料学 merges。朴素实现每轮全语料重数所有 pair 频次——词表 5 万意味着数万轮 × 全语料扫描,不可接受。
两个问题的共同解法:用对数据结构,让每轮只动「变了的那一点」。
2. 直觉:链表管顺序,堆管优先级
2.1 推理侧
把词想成一列火车车厢:
- 车厢是符号(初始每个字符一节),车厢之间用挂钩连(prev/next 指针)。合并两节车厢 = 把左边那节换成新符号、右边那节摘掉——不需要移动其他任何车厢。
- 所有「可合并的相邻对」放进一个按 rank 排队的优先队列。每次弹出 rank 最小的对;合并后只有「新符号与左右邻居」这两个新对可能进队。
- 堆里会有「过期票」(对的一侧早已被别的合并吃掉)——不主动清理,取出来时发现对不上就扔掉(惰性失效)。
2.2 训练侧
训练是镜像问题:每轮要全语料里最高频的 pair。窍门是合并某对只会改变「包含该对的词」内部的相邻关系——所以维护一个「pair → 出现在哪些词里」的索引(where_to_update),每轮只重算这些词贡献的计数差量,其余计数原封不动。
3. 图示:merge_all 的一轮
以词 u n d e r、merges 表 ("u","n")→rank 3、("n","d")→rank 1、("un","d")→rank 5… 为例:
初始: [u]<->[n]<->[d]<->[e]<->[r]
堆: {rank1:(n,d), rank3:(u,n)}
① 弹出 (n,d) rank1 ── 校验: n 的下一个确实是 d、且 merges[(n,d)] 的新 id 一致
合并: [u]<->[nd]<->[e]<->[r] (d 那节 len=0,等收尾统一清)
新对入堆: (u,nd) 查表有 rank5 → {rank5:(u,nd), rank3:(u,n)过期票}
② 弹出 (u,n) rank3 ── 校验: u 的下一个是 nd 不是 n → 过期,丢弃
③ 弹出 (u,nd) rank5 ── 校验通过 → 合并: [und]<->[e]<->[r]
新对 (und,e) 查表无 → 不入堆
④ 堆空 → retain(len>0) 清掉被摘车厢 → [und, e, r]
4. 原理演示(示意代码)
# 示意,非源码
def merge_all(chars, merges): # merges: (a,b) -> (rank, new_id)
syms = linked_list(chars)
heap = PriorityQueue()
for i, (a, b) in enumerate(adjacent_pairs(syms)):
if (a, b) in merges:
heap.push(rank=merges[(a,b)].rank, pos=i)
while heap:
top = heap.pop_min_rank()
if stale(top, syms, merges): # 惰性失效
continue
merge_at(top.pos) # 左车厢换新符号,右车厢摘除
for neighbor_pair in [left_new, new_right]:
if neighbor_pair in merges:
heap.push(neighbor_pair)
return [s for s in syms if s.alive]
真实实现与这份伪代码几乎一一对应,差别只在生产细节(dropout、过期校验的精确定义、收尾 retain)。
5. 真实实现(推理侧)
5.1 Symbol 与 Word:数组模拟的双向链表
Symbol(tokenizers/src/models/bpe/word.rs:39-44)四个字段:c(token id)、prev/next(数组下标,-1 表示端点)、len(该符号覆盖的字节数,摘除时置 0)。Word(同文件 57-59 行)就是 Vec<Symbol>。
合并两符号靠 merge_with(同文件 48-53 行):
// 摘自 tokenizers/src/models/bpe/word.rs:48-53
pub fn merge_with(&mut self, other: &Self, new_c: u32) {
self.c = new_c;
self.len += other.len;
self.next = other.next;
}
len 字段是 offset 的来源:最终每个符号的 (start, end) 由前缀和推出,BPE 全程不需要碰字符串本身。
5.2 merge_all:最小堆 + 惰性失效
Word::merge_all(tokenizers/src/models/bpe/word.rs:163-251)是推理核心。用的是四元堆 QuaternaryHeap(同文件 3 行 import),堆元素 Merge { pos, rank, new_id } 的比较器按 (rank, pos) 排序(同文件 15-36 行手写的 Ord——rank 小者优先,同 rank 位置靠左优先)。
主体循环(同文件 184-246 行)做的事与 §3 图完全对应:
// 摘自 tokenizers/src/models/bpe/word.rs:198-215(注释精简)
if self.symbols[top.pos].len == 0 { continue; } // 已被摘除
if self.symbols[top.pos].next == -1 { continue; } // 已是末尾
let next_pos = self.symbols[top.pos].next as usize;
let right = self.symbols[next_pos];
// 过期校验:当前这对在 merges 表里的 new_id 与入堆时不同 → 丢弃
let target_new_pair = (self.symbols[top.pos].c, right.c);
if merges.get(&target_new_pair).is_none_or(|(_, new_id)| *new_id != top.new_id) {
continue;
}
self.symbols[top.pos].merge_with(&right, top.new_id);
self.symbols[next_pos].len = 0; // 摘除右符号
随后把「左邻居 + 新符号」「新符号 + 右邻居」两个新对查表入堆(同文件 218-245 行),循环结束后 retain(|s| s.len != 0) 统一清理(251 行)。
复杂度: 每个符号最多进出堆常数次,O(n log n),且没有字符串重建。
5.3 dropout:随机跳过,训练正则化用的
循环开头有一个 dropout 分支(同文件 185-189 行):以概率 d 把弹出的 merge 塞进 skip 暂存,本轮先不合;一旦有一个 merge 真被执行,就把暂存的全放回堆(queue.extend(skip.drain(..)),187 行)。效果是「这次切分里某些合并被随机跳过 」——BPE-Dropout 论文的实现,推理端若设了 dropout 会走无缓存路径(tokenizers/src/models/bpe/model.rs:607-610)。
5.4 merge_word:词 → 初始符号序列
BPE::merge_word(tokenizers/src/models/bpe/model.rs:469-553)负责把字符串变成初始 Word:
- 逐字符查词表;非首字符拼
continuing_subword_prefix(如 WordPiece 的##)、末字符拼end_of_word_suffix(如 GPT-2 风格词表里的</w>)后再查(同文件 485-494 行); - 查不到时分三种回退:
byte_fallback把每字节映成<0xNN>再查(同文件 501-517 行);否则用unk_token,且fuse_unk决定连续未知片段是合成一个 unk 还是逐段 unk(同文件 519-545 行); - 最后
word.merge_all(&self.merges, self.dropout)(同文件 549 行)。
5.5 缓存:整词结果直接复用
tokenize_with_cache(tokenizers/src/models/bpe/model.rs:559-589)在 merge_word 外面包了一层「词 → Word」的缓存:线程本地、按模型实例隔离(BPE_LOCAL_CACHE + cache.id()),只缓存短于 MAX_LENGTH = 256 的词(tokenizers/src/utils/cache.rs:10),容量可配。自然文本里词频服从 Zipf 分布,这一层命中率高得惊人——高频词的 BPE 实际成本趋近于一次哈希查找。
入口 Model for BPE::tokenize(tokenizers/src/models/bpe/model.rs:601-612)只有三个分支:空串直接返回;有 dropout 走无缓存;否则走缓存路径。另有 ignore_merges 模式(同文件 563-570 行):整词直接查表,完全跳过合并——给「词表即全部」的退化用法。
6. 真实实现(训练侧)
6.1 feed:先复用管线切词计数
训练第一步不是 BPE 自己的事:BpeTrainer::feed(tokenizers/src/models/bpe/trainer.rs:645-670)拿 tokenizer 的 normalizer + pre_tokenizer 把语料切成词,统计 词 → 次数 的 AHashMap(并行用 maybe_par_bridge)。所以训练出的 merges 天然与推理时的预分词一致——这是「训推一致」最容易被忽略的来源。
6.2 do_train:五步主线
do_train(tokenizers/src/models/bpe/trainer.rs:456-624):
- 特殊 token 先入词表(
add_special_tokens); - 初始字母表:统计所有词的首字符(或全部字符,取决于配置)入表(
compute_alphabet); tokenize_words:把每个词表示成「符号 id 序列」的Word;count_pairs:统计全部相邻 pair 频次,并建立where_to_update: pair → {词下标集合};- 主循环(同文件 508-603 行):大顶堆
OctonaryHeap弹最高频 pair → 造新 token 入词表、记一条 merge → 只重算受影响词的 pair 差量 → 新 pair 入堆。终止条件:词表满 / 堆空 / 最高频低于min_frequency。
循环体里「过期条目」的处理和推理侧同款(同文件 517-521 行):弹出来的 count 与当前真实计数不符就刷新后重插一次,下次再弹出时若还是旧值自然被 min_frequency 或空堆终止掉。
6.3 增量更新 + 那个著名的 unsafe
合并一对后只需更新「包含该对的词」。top.pos 就是 where_to_update 给的词下标集合,更新循环(同文件 553-578 行)用了全库唯一一处刻意的 unsafe:
// 摘自 tokenizers/src/models/bpe/trainer.rs:557-578(注释精简)
struct WordPtr(*mut Word);
// Safety: 不做同内存并发访问,只访问同一分配内不同 chunk
unsafe impl Sync for WordPtr {}
let word_start = WordPtr(words.as_mut_ptr());
let changes = pos.maybe_par_iter().flat_map(|&i| {
unsafe {
let word = word_start.0.add(i);
(*word).merge(top.pair.0, top.pair.1, new_token_id, max_token_length)
.into_iter().map(|c| (c, i)).collect::<Vec<_>>()
}
}).collect::<Vec<_>>();
安全性的依据写在注释里:下标集合是 AHashSet<usize>,每个词只出现一次,所以并行写不同下标的 Word 不产生数据竞争;用裸指针是为了绕开「&mut 别名即 UB」的规则。Word::merge(tokenizers/src/models/bpe/word.rs:107-160)返回的是 Vec<(Pair, i32)>——该词里每个 pair 的计数差量(新对 +1、被拆掉的旧对 -1),主循环据此增量更新全局 pair_counts 并把新对入堆(同文件 580-597 行)。
6.4 灌回模型
训练结束后把 word_to_id、merges(pair → (rank, new_id),rank 即学习顺序)写回 BPE(同文件 606-623 行),continuing_subword_prefix 等选项一并转移。vocab 和 merges 从此就是数据,不是代码——save 出来就是 tokenizer.json 里那两个字段。
7. 坑与注意点
merges.txt的行序就是 rank。 加载时按行号赋 rank(tokenizers/src/models/bpe/model.rs:369-387的convert_merges_to_hashmap:lines.enumerate()的行号直接成为 rank);手写 merges 文件顺序错了,切分结果就全错。- byte_fallback 要求词表里有全部 256 个
<0xNN>。 缺一个字节,整个回退分支返回None然后落入 unk 逻辑(tokenizers/src/models/bpe/model.rs:501-517)。 - 缓存与 dropout 互斥:开了 dropout 必走无缓存路径(
tokenizers/src/models/bpe/model.rs:607-610),否则「随机性」会被缓存吃掉。 - 缓存不跨模型实例共享,且长词不缓存(≥256 字节,
tokenizers/src/utils/cache.rs:10)。超长输入(整段粘贴的 base64 之类)会稳定地慢。 min_frequency的终止是「硬停」:堆顶低于阈值直接 break(tokenizers/src/models/bpe/trainer.rs:524-526),即使词表还没满。- 词表满员的判定在循环开头(同文件 511-514 行)——特殊 token 和字母表占的名额会压缩实际可学的 merge 数。
- 训练时的
max_token_length在Word::merge里生效(tokenizers/src/models/bpe/word.rs:150-155的长度检查),超限的组合不会作为新对计入差量。 - Unigram 推理不是 merge 而是 Viterbi:在 lattice 上找最短 路径(
tokenizers/src/models/unigram/model.rs:347-352,lattice.viterbi()),还支持sample采样切分——别用 BPE 的心智模型套它。
8. 代码地图
| 主题 | 文件路径 | 符号名 |
|---|---|---|
| 推理核心 | tokenizers/src/models/bpe/word.rs | Word、Symbol、merge_all、merge、merge_with |
| 词 → 符号序列 | tokenizers/src/models/bpe/model.rs | BPE::merge_word、word_to_tokens |
| 缓存 | tokenizers/src/models/bpe/model.rs、tokenizers/src/utils/cache.rs | tokenize_with_cache、Cache、MAX_LENGTH |
| 模型入口 | tokenizers/src/models/bpe/model.rs | impl Model for BPE::tokenize |
| vocab/merges 加载 | tokenizers/src/models/bpe/model.rs | BPE::read_file、convert_merges_to_hashmap |
| 训练主线 | tokenizers/src/models/bpe/trainer.rs | BpeTrainer::do_train、count_pairs、tokenize_words、compute_alphabet |
| 训练喂数据 | tokenizers/src/models/bpe/trainer.rs | Trainer::feed |
| WordPiece | tokenizers/src/models/wordpiece/mod.rs | WordPiece::tokenize(最长前缀匹配) |
| Unigram | tokenizers/src/models/unigram/model.rs、lattice.rs | Unigram::tokenize、Lattice::viterbi、sample |
| WordLevel | tokenizers/src/models/wordlevel/mod.rs | WordLevel::tokenize |