推理阶段的BPE BPE 推理阶段详解如何正确编码与解码在BPE中三个核心功能为统计、合并与训练循环从原始文本中统计相邻字节对出现的频率不断合并最高频的字节对最终得到一份“合并规则表”merges和一份“词表”vocab。训练完成后我们就进入了推理阶段。推理阶段要解决的核心问题是给定一段新的文本如何利用训练好的 merges 规则把它转换成 Token ID 序列反过来给定一个 Token ID 序列如何还原成人类可读的文本这两个过程分别叫做编码encode和解码decode。1. 推理阶段编码的核心思想在训练 BPE 时我们是这样做的统计训练语料中所有相邻字节对的出现频率找出频率最高的那一对合并它们重复上述过程直到达到预设的合并次数或词表大小。这个过程中每一次合并都会产生一个新的 Token并被记录到merges表中。例如merges{(101,32):256,# 第 1 次合并产生 Token 256(116,104):257,# 第 2 次合并产生 Token 257(256,257):258,# 第 3 次合并产生 Token 258...}注意字典的value 是合并顺序号也就是新 Token 的 ID。这个顺序非常重要ID 越小说明这个合并发生得越早它在训练集中的优先级越高。那么在推理时我们面对一段新文本应该怎么合并呢核心原则严格遵循训练时确定的优先级顺序而不是重新统计当前文本中的频率。换句话说训练时我们“投票”选出了哪些字节对最值得合并并给它们排了序推理时我们不能再重新投票而是必须按照已经排好的顺序依次应用这些合并规则。很多初学者会犯一个错误在推理时仍然每次寻找当前文本中出现频率最高的字节对进行合并。这是完全错误的。因为推理时文本很短频率统计没有统计意义而且会破坏与训练时的一致性。2.encode函数逐行解析下面我们来看一个朴素但正确的encode实现。defencode(text,merges):# 1. 转为字节流idslist(text.encode(utf-8))whilelen(ids)2:# 获取当前文本中所有相邻对及其出现次数statsget_stats(ids)# 寻找“在 merges 规则表中存在且 ID 最小即最早被训练出来”的对pair_to_mergeNonemin_rankfloat(inf)# rank 即 IDforpairinstats:ifpairinmerges:rankmerges[pair]ifrankmin_rank:min_rankrank pair_to_mergepair# 如果当前序列中没有任何对在我们的规则表中停止ifpair_to_mergeisNone:break# 执行合并idsmerge(ids,pair_to_merge,min_rank)returnids2.1 第一步把文本转换为字节流idslist(text.encode(utf-8))BPE 的底层操作单位是字节byte而不是字符。这样做的好处是可以覆盖所有 Unicode 字符不会出现 OOV未登录词问题一个字符可能由多个字节组成例如中文“你”在 UTF-8 下是b\xe4\xbd\xa0对应三个字节[228, 189, 160]。因此编码的第一步就是把原始字符串编码成 UTF-8 字节序列然后把每个字节作为一个初始 Token ID。此时ids的长度就是文本的字节数。2.2 第二步进入合并循环whilelen(ids)2:只要序列中还有至少两个 Token就可能存在可以合并的相邻对。如果长度小于 2自然无法合并循环结束。2.3 第三步统计当前所有相邻对statsget_stats(ids)get_stats函数的作用是扫描整个ids列表统计所有相邻对pair的出现次数。例如ids[1,2,3,1,2]stats{(1,2):2,(2,3):1,(3,1):1}这里统计频率并不是为了选最高频的对而是为了知道当前有哪些相邻对存在。因为如果某个 pair 根本不存在于当前序列中我们就不需要考虑它。2.4 第四步寻找优先级最高的可合并对pair_to_mergeNonemin_rankfloat(inf)forpairinstats:ifpairinmerges:rankmerges[pair]ifrankmin_rank:min_rankrank pair_to_mergepair这一步是整个编码过程的核心。我们遍历当前所有出现的相邻对检查它们是否在merges规则表中。如果存在就取出它的rank即训练时的合并顺序 ID。由于rank越小表示越早被训练出来优先级越高所以我们用min_rank记录当前找到的最小 rank并记录对应的pair_to_merge。注意我们找的是“当前所有可合并对中 rank 最小的那个”而不是“当前出现频率最高的那个”。2.5 第五步如果没有可合并的对停止ifpair_to_mergeisNone:break如果遍历完所有相邻对发现没有任何一对出现在merges中说明当前序列已经无法再按照训练好的规则进行合并了。此时编码结束直接返回当前的ids。2.6 第六步执行合并idsmerge(ids,pair_to_merge,min_rank)merge函数会扫描整个ids列表将其中所有连续出现的pair_to_merge即(first, second)替换为新的 Token IDmin_rank。例如假设pair_to_merge (101, 32)min_rank 256那么merge([101,32,101,32,99],(101,32),256)# 返回 [256, 256, 99]注意一次合并会替换所有匹配的相邻对而不是只替换第一个。这与训练时的行为一致因为在训练时一次合并也会同时替换语料中所有该字节对的出现。合并完成后我们回到循环开头重新统计新的相邻对再次寻找下一个可合并的最小 rank 对。如此反复直到没有可合并的 pair 为止。3. 为什么必须按 rank 顺序一个反例很多同学会问为什么不能每次找当前文本中出现频率最高的对或者为什么不能随便选一个可合并的对我们来看一个具体的例子。假设训练后我们得到了以下两条规则merges{(97,98):256,# rank 1: a b - X(98,99):257,# rank 2: b c - Y}现在输入文本是abc对应的字节序列是[97, 98, 99]。正确的编码过程按 rank 顺序当前相邻对(97, 98)和(98, 99)。两者都在merges中rank 分别是 1 和 2。选择 rank 最小的(97, 98)合并成 256得到[256, 99]。此时相邻对只有(256, 99)它不在merges中停止。最终 Token ID 序列为[256, 99]。错误的编码过程比如先合并 rank 2 的(98, 99)当前相邻对(97, 98)和(98, 99)。如果我们不按 rank而是随便选了(98, 99)合并成 257得到[97, 257]。此时相邻对是(97, 257)它不在merges中停止。最终 Token ID 序列为[97, 257]。可以看到两种顺序得到了完全不同的 Token 序列。为什么必须选 rank 最小的因为 BPE 的训练过程是逐步、贪心的。在训练时我们第一步合并的是(a, b)因为它频率最高。合并之后语料中出现了新的 TokenX然后第二步我们才合并(b, c)。注意当我们合并(b, c)时语料中可能已经存在X但X后面可能跟着c而b和c的直接相邻情况可能已经减少了。推理时如果我们不按训练时的顺序来就可能造成合并的上下文不一致。例如在训练时(a, b)合并成X后X和c的组合可能没有被后续合并规则覆盖但在推理时如果先合并(b, c)成Y那么a和Y的组合可能完全不在训练时出现过。这会导致生成的 Token 序列在训练分布中从未存在过模型无法正确理解。简单说训练时的顺序定义了“什么先合并、什么后合并”的优先级。推理时必须遵守这个优先级才能保证编码结果与训练时的语言模型分布一致。4. 优化与工程实现上面给出的encode实现是朴素版本时间复杂度较高。因为每次合并后都要重新扫描整个序列统计相邻对而一个长文本可能包含几十万个字节合并次数也可能很多。在实际工程中我们会做以下优化4.1 使用优先队列最小堆我们可以维护一个优先队列里面存放当前所有相邻对及其对应的 rank。每次从队列中取出 rank 最小的 pair 进行合并然后只更新受影响的相邻对而不是重新扫描整个序列。例如在minbpe的BPEEncoder实现中就使用了类似的思想维护一个heap来动态获取当前最小 rank 的 pair合并后只更新局部信息。4.2 正则表达式分块处理在 GPT-2 / GPT-4 等模型中BPE 并不是直接对整个文本进行无脑合并而是先用正则表达式把文本切成一个个“块”chunk然后在每个块内部分别进行 BPE 合并。这样做有两个好处减少序列长度长文本切成小块后每个块的处理更快。避免跨边界合并某些标点、空格、换行等不应该与相邻的单词合并正则分块可以强制这些边界不被跨越。例如GPT-2 使用的正则表达式大致是patternrs|t|re|ve|m|ll|d| ?\p{L}| ?\p{N}| ?[^\s\p{L}\p{N}]|\s(?!\S)|\s它会匹配出单词、数字、标点、空格等不同的块然后对每个块分别进行 BPE 编码。这样编码结果更加稳定也符合语言习惯。5.decode函数详解解码是编码的逆过程逻辑上简单很多但有一个非常关键的细节如何处理无效的 UTF-8 字节序列。5.1 构建vocab在解码时我们需要一个vocab字典它把 Token ID 映射回它对应的字节串bytes。这个vocab是怎么来的呢其实在训练 BPE 的过程中我们就可以顺便构建它# 初始化基础字节 Token 0~255 对应单个字节vocab{i:bytes([i])foriinrange(256)}# 每次合并时记录新 Token 对应的字节串for(p1,p2),new_idinmerges.items():vocab[new_id]vocab[p1]vocab[p2]这样每个 Token ID 都能唯一映射到一个字节串。例如vocab[97]→bavocab[256]→vocab[97] vocab[98]→babvocab[258]→vocab[256] vocab[257]→bab bbc→babbc有了vocab解码就非常简单了。5.2 解码函数defdecode(ids,vocab): ids: token ID 列表 vocab: 映射 {idx: bytes} (由基础字节和 merges 反推得到) # 将所有 ID 映射回字节串并拼接tokensb.join(vocab[idx]foridxinids)# errorsreplace 是关键texttokens.decode(utf-8,errorsreplace)returntext第一步把每个 Token ID 转换成对应的 bytes然后拼接成一个完整的字节串。第二步把这个字节串解码成 UTF-8 字符串。5.3 为什么需要errorsreplace这是很多初学者容易忽略的细节。在 Python 中如果我们直接对一个包含非法 UTF-8 字节序列的字节串调用.decode(utf-8)程序会抛出UnicodeDecodeError异常。例如b\xe4\xbd.decode(utf-8)# 会抛出 UnicodeDecodeError为什么会出现非法 UTF-8 字节序列原因在于BPE 是在字节级别进行合并的它并不保证每个 Token 对应的字节串都是合法的 UTF-8 字符。更常见的情况是模型生成的 Token 序列被截断了。例如一个中文字符“你”在 UTF-8 下由三个字节组成b\xe4\xbd\xa0。如果模型因为max_tokens限制只输出了前两个字节对应的 Token ID然后停止了那么解码时我们就会得到b\xe4\xbd。这显然不是一个合法的 UTF-8 字符。如果我们不使用errorsreplace程序就会崩溃。使用errorsreplace后Python 会把无法解码的字节序列替换成 Unicode 替换字符UFFFD。这样程序不会崩溃用户也能看到输出中有一个占位符提示这里出现了截断或无效字节。5.4 一个具体的例子假设我们的vocab中有以下 Tokenvocab{256:b\xe4\xbd,# 不完整的“你”的前两个字节257:b\xa0,# “你”的最后一个字节}如果模型生成了 Token ID 序列[256, 257]解码过程如下tokensb\xe4\xbdb\xa0b\xe4\xbd\xa0texttokens.decode(utf-8,errorsreplace)# text 你这是正常情况。但如果模型只生成了[256]那么tokensb\xe4\xbdtexttokens.decode(utf-8,errorsreplace)# text 或可能是两个替换字符取决于具体实现这样程序就不会崩溃而是输出了一个替换字符。在 GPT-3、GPT-4 的输出中我们偶尔会看到就是这个原因。6. 完整可运行示例下面给出一个完整的 BPE 编码与解码的朴素实现方便大家理解整个流程。fromcollectionsimportCounterdefget_stats(ids):统计相邻对出现次数returnCounter(zip(ids,ids[1:]))defmerge(ids,pair,new_id):将 ids 中所有连续出现的 pair 替换为 new_idnew_ids[]i0whileilen(ids):ifilen(ids)-1andids[i]pair[0]andids[i1]pair[1]:new_ids.append(new_id)i2else:new_ids.append(ids[i])i1returnnew_idsdefencode(text,merges):idslist(text.encode(utf-8))whilelen(ids)2:statsget_stats(ids)pair_to_mergeNonemin_rankfloat(inf)forpairinstats:ifpairinmerges:rankmerges[pair]ifrankmin_rank:min_rankrank pair_to_mergepairifpair_to_mergeisNone:breakidsmerge(ids,pair_to_merge,min_rank)returnidsdefbuild_vocab(merges):根据基础字节和 merges 构建 vocabvocab{i:bytes([i])foriinrange(256)}for(p1,p2),new_idinmerges.items():vocab[new_id]vocab[p1]vocab[p2]returnvocabdefdecode(ids,vocab):tokensb.join(vocab[idx]foridxinids)returntokens.decode(utf-8,errorsreplace)# 示例 merges训练时得到的规则merges{(97,98):256,# a b - 256(98,99):257,# b c - 257}vocabbuild_vocab(merges)textabcidsencode(text,merges)print(Token IDs:,ids)# 输出: [256, 99]print(Decoded:,decode(ids,vocab))# 输出: abc运行这个示例你会发现encode(abc)得到的是[256, 99]而不是[97, 257]。这正是因为我们遵循了 rank 顺序。