跳到主要内容

数据截至 (上游 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:数组模拟的双向链表

Symboltokenizers/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_alltokenizers/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_wordtokenizers/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_cachetokenizers/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::tokenizetokenizers/src/models/bpe/model.rs:601-612)只有三个分支:空串直接返回;有 dropout 走无缓存;否则走缓存路径。另有 ignore_merges 模式(同文件 563-570 行):整词直接查表,完全跳过合并——给「词表即全部」的退化用法。


6. 真实实现(训练侧)

6.1 feed:先复用管线切词计数

训练第一步不是 BPE 自己的事:BpeTrainer::feedtokenizers/src/models/bpe/trainer.rs:645-670)拿 tokenizer 的 normalizer + pre_tokenizer 把语料切成词,统计 词 → 次数AHashMap(并行用 maybe_par_bridge)。所以训练出的 merges 天然与推理时的预分词一致——这是「训推一致」最容易被忽略的来源。

6.2 do_train:五步主线

do_traintokenizers/src/models/bpe/trainer.rs:456-624):

  1. 特殊 token 先入词表(add_special_tokens);
  2. 初始字母表:统计所有词的首字符(或全部字符,取决于配置)入表(compute_alphabet);
  3. tokenize_words:把每个词表示成「符号 id 序列」的 Word
  4. count_pairs:统计全部相邻 pair 频次,并建立 where_to_update: pair → {词下标集合}
  5. 主循环(同文件 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::mergetokenizers/src/models/bpe/word.rs:107-160)返回的是 Vec<(Pair, i32)>——该词里每个 pair 的计数差量(新对 +1、被拆掉的旧对 -1),主循环据此增量更新全局 pair_counts 并把新对入堆(同文件 580-597 行)。

6.4 灌回模型

训练结束后把 word_to_idmerges(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-387convert_merges_to_hashmaplines.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_lengthWord::merge 里生效tokenizers/src/models/bpe/word.rs:150-155 的长度检查),超限的组合不会作为新对计入差量。
  • Unigram 推理不是 merge 而是 Viterbi:在 lattice 上找最短路径(tokenizers/src/models/unigram/model.rs:347-352lattice.viterbi()),还支持 sample 采样切分——别用 BPE 的心智模型套它。

8. 代码地图

主题文件路径符号名
推理核心tokenizers/src/models/bpe/word.rsWordSymbolmerge_allmergemerge_with
词 → 符号序列tokenizers/src/models/bpe/model.rsBPE::merge_wordword_to_tokens
缓存tokenizers/src/models/bpe/model.rstokenizers/src/utils/cache.rstokenize_with_cacheCacheMAX_LENGTH
模型入口tokenizers/src/models/bpe/model.rsimpl Model for BPE::tokenize
vocab/merges 加载tokenizers/src/models/bpe/model.rsBPE::read_fileconvert_merges_to_hashmap
训练主线tokenizers/src/models/bpe/trainer.rsBpeTrainer::do_traincount_pairstokenize_wordscompute_alphabet
训练喂数据tokenizers/src/models/bpe/trainer.rsTrainer::feed
WordPiecetokenizers/src/models/wordpiece/mod.rsWordPiece::tokenize(最长前缀匹配)
Unigramtokenizers/src/models/unigram/model.rslattice.rsUnigram::tokenizeLattice::viterbisample
WordLeveltokenizers/src/models/wordlevel/mod.rsWordLevel::tokenize