9.BPE Tokenizer 训练伪代码
从文本语料中统计所有单词及其频率,输出词频字典¶
函数目标:输入一段原始文本字符串,输出一个字典,键为单词( token ),值为该单词在语料中出现的次数。这里假设单词由空格分隔,并且我们希望保留大小写与标点符号,通常 BPE 训练前会进行基础分词(如按空格分词),同时在每个单词末尾保留词边界信息,因此统计的是“单词”粒度,后续会进一步拆成字符并加结束符。
实现思路:
-
将文本按空白字符(空格、换行、制表符等)切分成单词列表。
-
遍历列表,使用字典统计出现次数。
-
可以选择是否进行简单的标点处理:如果对标点有特殊要求,例如“word,”和“word”被视为不同单词,则直接保留;若需要分离标点,可先做分词预处理(如用正则表达式将标点与单词分开)。这里给出基础版本:按空格分词,保留原样。
-
考虑到大规模语料,可以使用
collections.Counter提高效率。
Python 代码示例:
from collections import Counter
import re
def build_word_freq(text: str) -> dict:
"""
统计文本中每个单词的出现频率。
参数:
text: 原始文本字符串。
返回:
word_freq: {word: freq} 字典。
"""
# 使用正则表达式按空白字符分割,同时保留单词与标点的连接?
# 这里简单按空格分割,适用于已预处理过的文本。
# 若文本中包含换行等,split() 默认按任意空白字符分割,且会去掉空字符串。
words = text.split()
word_freq = Counter(words)
return dict(word_freq)
注意事项:
-
如果希望将标点与单词分开(如 "hello." 变成 "hello", "."),可在统计前进行分词,例如使用
re.findall(r'\b\w+\b|\S', text)或更复杂的 tokenizer。 -
真实的 BPE 实现(如 GPT 系列)在训练前会先将文本按字节或字符进行初始分割,但词频统计仍是以空格分隔的单词为单位,因为 BPE 需要知道每个单词的频率,以便在单词内部进行子词合并。
示例:
text = "low lower lowest low low"
word_freq = build_word_freq(text)
print(word_freq) # {'low': 3, 'lower': 1, 'lowest': 1}
初始化词汇表:将单词拆分为字符序列,加结束符 </w>,统计字符初始频率¶
目标:给定词频字典 word_freq,将所有单词转化为字符序列(每个字符作为一个符号),并在序列末尾添加特殊符号 </w> 表示单词结束。同时统计每个初始字符在整个语料中的频率(频率 = 单词频率 × 该字符在该单词中的出现次数)。
实现细节:
-
遍历
word_freq的每一项(word, freq)。 -
将单词转换为字符列表,例如
"low"->['l', 'o', 'w'],然后追加'</w>'。 -
存储每个单词对应的符号序列,通常用一个字典
word2symbols保存,键为原始单词,值为符号列表;同时为了后续合并,我们还需要知道每个单词的频率,因此可以存为元组(symbol_list, freq)。 -
字符频率统计:对于每个单词的每个符号(包括
</w>),累加freq次。可以用Counter或普通字典。
Python 代码:
def initialize_vocab(word_freq: dict):
"""
初始化词汇表,返回单词符号序列列表和字符频率字典。
参数:
word_freq: {word: freq}
返回:
word_symbols: 列表,元素为 (symbol_list, freq)
char_freq: {symbol: total_freq}
"""
word_symbols = [] # 存储每个单词的符号序列及其频率
char_freq = Counter()
for word, freq in word_freq.items():
symbols = list(word) + ['</w>']
word_symbols.append((symbols, freq))
for sym in symbols:
char_freq[sym] += freq
return word_symbols, dict(char_freq)
说明:
-
结束符
</w>的作用是标记单词边界。在解码时,遇到</w>就知道这里是词尾,后面应该加上空格(除非是最后一个词)。这避免了合并时跨单词边界产生无意义的子词。 -
初始化后,词汇表里的基本符号就是所有出现的字符加上
</w>。
示例:
word_freq = {"low": 3, "lower": 1, "lowest": 1}
word_symbols, char_freq = initialize_vocab(word_freq)
# word_symbols:
# [(['l','o','w','</w>'], 3),
# (['l','o','w','e','r','</w>'], 1),
# (['l','o','w','e','s','t','</w>'], 1)]
# char_freq: {'l':5, 'o':5, 'w':5, '</w>':5, 'e':2, 'r':1, 's':1, 't':1}
计算所有相邻符号对(bigram)频率¶
目标:给定带频率的单词符号序列列表(如上一问的 word_symbols),统计所有相邻符号对的出现频次。频率的计算方式:对于某个单词,如果它的符号序列中有相邻对 (sym_i, sym_{i+1}),则该 bigram 获得该单词的频率次数(因为该单词在语料中出现了 freq 次,每次出现都会贡献一次这个相邻对)。因此 bigram 频率 = 所有包含该 bigram 的单词的频率之和。
实现:
-
遍历
word_symbols中的每个(symbols, freq)。 -
在
symbols序列中,从索引 0 到 len-2,依次取出(symbols[i], symbols[i+1])。 -
将 pair 加入字典,累加
freq。
代码:
def compute_bigram_freq(word_symbols: list) -> dict:
"""
计算所有 bigram 的频率。
参数:
word_symbols: [(symbol_list, freq), ...]
返回:
bigram_freq: {(sym1, sym2): freq}
"""
bigram_freq = Counter()
for symbols, freq in word_symbols:
for i in range(len(symbols) - 1):
pair = (symbols[i], symbols[i+1])
bigram_freq[pair] += freq
return dict(bigram_freq)
注意:
-
如果某个单词频率很高,其内部 bigram 会被多次累加,这正确反映了语料中的统计特性。
-
在 BPE 训练过程中,每次合并后会改变某些单词的符号序列,需要重新计算 bigram 频率。为了提高效率,实际实现中会在合并后局部更新 bigram 频率,而不是每次全量重算。
示例:
基于上述 word_symbols:
-
('l','o')出现于三个单词,频率之和=3+1+1=5 -
('o','w')同理频率=5 -
('w','</w>')只在 "low" 中出现,频率=3 -
('w','e')在 "lower" 和 "lowest" 中出现,频率=1+1=2 等等。
单步 BPE 合并:找出频率最高的 bigram¶
目标:给定 bigram 频率字典,返回频率最高的符号对。如果存在多个频率相同的最高频 pair,则按字典序选择(例如比较元组在 Python 中的默认顺序)。
实现:
-
使用
max函数,指定 key 为频率,但需要处理频率相同的情况。 -
Python 的
max在比较元组时,若第一个元素(频率)相同,会比较第二个元素(key),即 bigram 元组本身,这正是字典序。
代码:
def select_best_pair(bigram_freq: dict) -> tuple:
"""
返回频率最高的 bigram,频率相同按字典序。
参数:
bigram_freq: {(sym1, sym2): freq}
返回:
best_pair: (sym1, sym2) 元组
"""
if not bigram_freq:
return None
# 用 max,比较 (freq, pair) 会导致按 freq 优先,freq 相同则按 pair 排序,
# 但我们需要的是先 freq 降序,再 pair 升序。所以可以自定义 key = lambda x: (x[1], x[0]) 但需要处理大小顺序。
# 简洁方式:先找到最大频率,再在最大频率的 pair 中取最小字典序。
max_freq = max(bigram_freq.values())
best_pair = min(pair for pair, freq in bigram_freq.items() if freq == max_freq)
return best_pair
说明:
-
返回的 best_pair 将作为下一步要合并的符号对,例如
('l', 'o')。 -
之所以规定频率相同按某种规则选择(如字典序),是为了保证算法的确定性,避免随机性。
示例:
若 bigram_freq = {('l','o'):5, ('o','w'):5},最高频率为5,相同频率下按字典序 ('l','o') < ('o','w'),因此选择 ('l','o')。
在词汇表中执行一次合并¶
目标:给定当前所有单词的符号序列 word_symbols 和选出的最高频 bigram best_pair = (a, b),将每个单词符号序列中所有相邻出现的 a, b 替换为新的合并符号 "ab"(通常新符号命名为 a+b 或直接拼接字符串)。更新符号序列,同时更新词汇表(即添加新符号,但实际词汇表可以从所有符号序列中收集)。
注意:
-
合并是“一次性”应用于所有单词的所有匹配位置,而不是只替换一次。
-
替换后可能会产生新的相邻对,例如原本序列
... X a b Y ...替换后变为... X ab Y ...,相邻对(X, ab)和(ab, Y)被创建,原本的(X, a),(a, b),(b, Y)会消失(除非其他地方还有)。 -
实现时通常直接修改
word_symbols列表,对每个单词的符号序列进行扫描与合并。
代码:
def merge_vocab(word_symbols: list, best_pair: tuple) -> list:
"""
在单词符号序列中合并指定的 bigram,返回更新后的 word_symbols。
参数:
word_symbols: [(symbol_list, freq), ...]
best_pair: (sym1, sym2) 待合并对
返回:
new_word_symbols: 更新后的列表(原地修改或新列表皆可,此处选择新建)
"""
a, b = best_pair
new_symbol = a + b # 新符号为两个字符串的拼接
new_word_symbols = []
for symbols, freq in word_symbols:
new_syms = []
i = 0
while i < len(symbols):
# 如果当前符号是 a 且下一个是 b,则替换为新符号
if i < len(symbols) - 1 and symbols[i] == a and symbols[i+1] == b:
new_syms.append(new_symbol)
i += 2
else:
new_syms.append(symbols[i])
i += 1
new_word_symbols.append((new_syms, freq))
return new_word_symbols
注意事项:
-
合并时,新符号的名字建议可读,如
"ab"。在真实场景中,可能会将 Unicode 字符直接拼接,若 a,b 本身可能已是多字符子词,拼接名字会很长,但逻辑相同。 -
更新后,需要重新计算 bigram 频率以进行下一次迭代;或者可以增量更新 bigram 频率表以提升效率。但在本框架中,步骤6会循环执行计算bigram->选最优->合并。
-
合并后的
new_word_symbols中,词汇表自动扩展,新增了new_symbol,原有符号可能不再出现。
示例:
原序列:['l','o','w','</w>'], best_pair=('l','o')
合并后:['lo','w','</w>']。
BPE 训练的完整伪代码¶
输入:原始文本 text,期望的合并次数 num_merges(或达到某个词汇表大小)。
输出:合并规则列表 merges(记录每次合并的 pair 和新符号),以及最终词汇表。
伪代码:
算法 BPE_Train(text, num_merges):
# 1. 词频统计
word_freq = 统计文本中每个单词的频率(按空格分词)
# 2. 初始化
word_symbols = [] # 元素为 (符号序列, 频率)
for each (word, freq) in word_freq:
symbols = list(word) + ['</w>']
word_symbols.append((symbols, freq))
vocab = 所有初始符号的集合(字符 + '</w>')
merges = [] # 记录合并规则,按顺序存储
# 3. 循环合并
for i from 1 to num_merges:
# 3.1 计算所有 bigram 频率
bigram_freq = {}
for each (symbols, freq) in word_symbols:
for j from 0 to len(symbols)-2:
pair = (symbols[j], symbols[j+1])
bigram_freq[pair] += freq
# 如果已经没有 bigram 可合并,提前结束
if bigram_freq 为空:
break
# 3.2 选择最高频 bigram
best_pair = 频率最高的 pair,若相同按字典序选取
# 3.3 记录合并规则(新符号 = best_pair[0] + best_pair[1])
new_symbol = best_pair[0] + best_pair[1]
merges.append( (best_pair, new_symbol) )
# 3.4 将合并应用到所有单词序列
new_word_symbols = []
for each (symbols, freq) in word_symbols:
new_syms = []
i = 0
while i < len(symbols):
if i < len(symbols)-1 and symbols[i] == best_pair[0] and symbols[i+1] == best_pair[1]:
new_syms.append(new_symbol)
i += 2
else:
new_syms.append(symbols[i])
i += 1
new_word_symbols.append((new_syms, freq))
word_symbols = new_word_symbols
# 3.5 更新词汇表(新符号加入)
vocab.add(new_symbol)
# 返回合并规则(可用于编码)和最终词汇表(所有符号)
return merges, vocab, word_symbols
说明:
-
返回的
merges是一个列表,顺序记录了每次合并的( (a,b), ab )。在编码时,将按照这个顺序尝试合并(从最早到最晚)。 -
有些 BPE 实现会在合并时保留原始符号并添加新符号,词汇表逐渐扩大;上述过程完全保留了最终所有符号的集合。
-
实际上,也可以通过优先队列(最大堆)来快速获取最高频 bigram,并在合并后只更新受影响的 bigram 频率,避免每轮重新扫描全部。但上述伪代码展示了最直观的逻辑。
编码函数:给定合并规则,将单词拆分为子词序列¶
目标:给定一个合并规则列表 merges(顺序即训练时的合并顺序)和一个原始单词字符串,输出该单词的子词序列(token 列表)。编码过程不需要单词频率,只需要按照学习到的合并规则对字符序列逐步合并。
标准做法:
-
将单词拆分为字符列表,末尾加上
</w>。 -
按照
merges的顺序,依次处理每条规则( (a,b), ab ): - 扫描当前符号序列,若发现相邻的
a后紧跟b,则将其合并为ab。 -
对于同一条规则,需要重复扫描直到没有可合并的为止吗?实际上,由于规则的顺序性,通常的做法是:按照 rules 的先后顺序,针对每一条规则,扫描序列并将所有满足条件的相邻对进行合并(一次性替换所有出现)。这与训练时的单步合并逻辑相同:训练时每一轮只合并一种 pair,但会对所有单词的所有出现进行合并;编码时也应该对所有出现进行合并,且按照规则顺序依次执行。因为后面的合并规则可能依赖于前面的合并结果,所以必须严格按顺序。
-
最终得到的符号序列就是该单词的子词划分。
实现细节:
-
由于合并规则很多,若对每个单词都从字符开始逐个应用所有规则,效率较低。但这是 BPE 编码的基本方法,实际优化(如使用缓存、前缀树等)可加速。
-
注意:在应用某条规则时,应从头到尾扫描,合并后可能立即形成新的可合并对,但规则规定只应用一次?训练时一次合并只针对特定的 pair,且在一次循环内,我们会把所有该 pair 都合并掉,但不会在合并后继续合并该 pair 的新出现(因为扫描是单次从左到右,如果合并后产生了新的相同 pair,但位于合并后的位置,是否再次合并?)。例如序列
a a a,规则('a','a') -> 'aa',从左到右扫描:索引0,1合并为'aa',此时序列变成['aa', 'a'],索引往后移,不会再合并新的('aa','a')因为当前规则是('a','a'),不会匹配('aa','a')。所以一次扫描只会合并非重叠的、原有的 pair。这与训练时的单步合并一致。因此我们实现时也采用同样的单趟扫描替换。
代码:
def encode_word(word: str, merges: list) -> list:
"""
使用学到的 BPE 合并规则对单词进行编码。
参数:
word: 原始单词字符串
merges: 列表,元素为 ((a,b), ab),按训练顺序
返回:
子词列表
"""
symbols = list(word) + ['</w>']
for (a, b), ab in merges:
# 应用规则 (a,b) -> ab
new_symbols = []
i = 0
while i < len(symbols):
if i < len(symbols) - 1 and symbols[i] == a and symbols[i+1] == b:
new_symbols.append(ab)
i += 2
else:
new_symbols.append(symbols[i])
i += 1
symbols = new_symbols
return symbols
注意事项:
-
编码函数应当与训练时的合并逻辑完全一致,才能保证编码后的子词序列可以通过解码还原(结合
</w>处理)。 -
返回的子词列表可能包含单个字符和合并后的多字符符号。
示例:
假设训练后 merges 顺序为:(('l','o'), 'lo'), (('lo','w'), 'low')。
单词 "low" -> 字符 ['l','o','w','</w>']。
应用第一条:合并 ('l','o') -> ['lo','w','</w>']。
应用第二条:合并 ('lo','w') -> ['low','</w>']。
最终子词序列:['low', '</w>']。
解码函数:将子词序列还原为原始字符串¶
目标:给定子词列表(即编码后的 token 序列)和合并规则(或词汇表),将其拼接成完整的原始文本。关键点在于正确处理结束符 </w>,它表示一个单词的结束,在还原后应当在对应位置添加空格(除了文本末尾的最后一个词)。
解码逻辑:
-
直接将所有子词符号拼接成一个长字符串。
-
然后将所有的
</w>替换为空格(或直接删除</w>并视作词边界)。 -
最后处理首尾多余的空格。对于文本最后一个单词,若其末尾有
</w>,替换后会多出一个尾部空格,通常可以去除末尾空格。
如果编码时是一整个单词的序列(即多个单词),解码时需要将子词拼接并还原出单词之间的空格。通常是整个文本的 token 列表,每个 token 可能是 'low', '</w>', 'er', '</w>' 等。解码时直接拼接得到 "low</w>er</w>",然后将 </w> 替换为空格得到 "low er ",再去掉首尾空格。注意 </w> 之前可能直接是子词,替换后单词间自然有空格。
实现:
def decode(tokens: list) -> str:
"""
将子词 token 列表解码为原始字符串。
参数:
tokens: 子词列表,包含 '</w>' 表示词尾
返回:
原始文本字符串
"""
# 拼接所有 token
text = ''.join(tokens)
# 将结束符替换为空格
text = text.replace('</w>', ' ')
# 去除首尾多余空格并规范化中间空格(多个空格变成一个,但一般不会出现)
text = ' '.join(text.split())
return text
说明:
-
如果 tokens 中不含
</w>,则意味着单词未结束,但正常编码会在每个单词末尾加上</w>,解码后能得到正确的空格分隔。 -
对于某些 BPE 实现,结束符可能是其他符号(如
_或特殊 token),处理方式类似:将结束符替换为空格。 -
如果解码的序列是整个文档的 token 序列,经过上述替换后就能完美复原原始文本(除了大小写等预处理保留的情况)。
示例:
tokens = ['low', '</w>', 'low', 'er', '</w>']
拼接 -> "low</w>lower</w>"
替换 </w> 为空格 -> "low lower "
规范化 -> "low lower"
扩展:
在真实的 NLP 流程中,</w> 只是一个特殊的词尾标记,解码时直接去除并添加空格即可。有些分词器会使用不同的策略,如将词首标记为 Ġ(GPT-2 中的空格符),解码时将 Ġ 转换为空格。但原理类似,都是通过特殊符号恢复原始空格边界。
处理未知词(OOV):在编码时,若出现词汇表中未见的字符,实现回退到字符级或使用 <UNK> 的策略¶
在 BPE 分词器中,未知词通常指包含训练时未见过字符的单词。由于 BPE 的初始词汇表建立在字符级,理想情况下如果训练语料覆盖了所有可能的 Unicode 字符,则不会出现未知字符。但实际中,测试文本可能包含罕见字符或表情符号等。针对这种情况,有两种主流策略:
-
回退到字符级(Fallback to characters):对未见字符仍然保留其原始字符作为 token,即进行字符级切分。因为即使某些字符没有在合并规则中出现过,但如果在初始化时将任意字符视为合法符号,那么编码时可直接将未知字符当作单字符 token 处理。这要求训练时采用了“开放词汇”策略,如 Byte-Level BPE(以字节为基础,256 个基本符号,可表示任何 UTF-8 序列),或训练时将词汇表初始化为所有 Unicode 码点(不现实,通常使用字节)。
-
使用
<UNK>特殊符号:训练时预留一个特殊 token<UNK>,编码时将所有未见过的字符替换为<UNK>。这会导致信息丢失,且解码后无法还原原始文本,通常不推荐用于生成任务,但在某些分类任务中可用。
由于前文我们采用字符级初始化,若测试时出现未见于初始字符集(或词汇表)的字符,编码过程会因找不到该字符而失败。因此需要扩展编码函数,加入 OOV 处理。更常见的做法是直接采用 Byte-Level BPE(见下一问),从根本上解决 OOV 问题,因为它将文本视为字节序列,任何文本都能用 256 个字节表示。
基于字符级 BPE 的 OOV 处理策略实现:
方案一:回退到字符本身(假设初始字符集为训练时见过的所有字符,未见字符仍保留原样,但无法参与任何合并,因为合并规则中不包含它们)。编码时,将单词拆分为字符,若某字符不在初始字符集中,则将其原样保留,且不参与合并(因为合并规则都以已知字符对为基础)。此时该未知字符会作为独立 token 输出。这种方法保持了文本的可还原性,只要解码时按原样拼接即可。但需要确保模型能处理这样的“扩展”词汇。
方案二:将未知字符替换为 <UNK>。在编码前预处理:检查每个字符是否在基础字符集中,若不在则替换为 <UNK>。<UNK> 在训练时作为特殊符号加入词汇表(频率可设为 0 或极低),编码时 <UNK> 可与其他符号合并吗?通常 <UNK> 不参与合并,它作为一个原子 token 使用。因此合并规则中可能不包含 <UNK> 相关的 pair。为了支持 <UNK>,在初始化时可以将一个虚拟单词 "<UNK>" 加入词频字典(频率为 0),从而让 < , U, N, K, > 和 </w> 出现在初始字符集中,之后可能被合并为一个整体 token <UNK>。但更简单的是将 <UNK> 作为单个不可分割的符号处理,编码时若遇到未知字符直接输出 ['<UNK>']。
下面给出结合两种策略的编码函数增强版:
def encode_word_with_oov(word: str, merges: list, known_chars: set,
unk_token: str = '<UNK>', use_unk: bool = False) -> list:
"""
编码一个单词,支持 OOV 字符处理。
参数:
word: 原始单词
merges: BPE 合并规则
known_chars: 训练时初始字符集(包含 '</w>')
unk_token: 未知词标记
use_unk: True 则使用 <UNK> 替代未知字符;False 则回退到字符本身
返回:
token 列表
"""
if use_unk:
# 将未知字符替换为 unk_token
chars = []
for ch in word:
if ch in known_chars:
chars.append(ch)
else:
chars.append(unk_token)
# 注意:若整个词变成 <UNK> 序列,仍按正常合并流程,但通常我们不合并 <UNK>
# 简单起见,直接返回 [unk_token] 如果包含任何未知字符? 可自定义策略。
# 这里我们保留 <UNK> 作为字符参与可能的合并(根据训练规则),但多数情况下直接返回 [unk_token] 表示整个词未知。
# 更稳健:若包含任何未知字符,直接输出 [unk_token] 并跳过合并?
# 以下提供按字符替换的版本:
else:
# 回退字符本身:未知字符保留原样,但不参与任何合并(因为规则中不包含它们)
chars = list(word)
# 后续编码同上,将 chars 加上 </w>
symbols = chars + ['</w>']
for (a, b), ab in merges:
new_symbols = []
i = 0
while i < len(symbols):
if i < len(symbols)-1 and symbols[i] == a and symbols[i+1] == b:
new_symbols.append(ab)
i += 2
else:
new_symbols.append(symbols[i])
i += 1
symbols = new_symbols
return symbols
然而,最彻底的解决方法是采用 Byte-Level BPE,如问题 10 所述。它将文本转为字节序列,初始符号就是 0~255 共 256 个字节。这样任何 Unicode 字符(包括表情符号)都被表示成多个字节,永远不会有未登录符号。GPT-2、GPT-3、LLaMA 等模型均采用此种方式。
实现 Byte-Level BPE 的初始化:输入文本,将其转换为 UTF-8 字节序列,并以字节作为初始符号,写出初始化逻辑¶
Byte-Level BPE 的核心思想是将每个单词视为 UTF-8 编码后的字节序列,每个字节(0-255)作为一个初始符号,而不是 Unicode 字符。这样做的好处是词汇表大小从海量 Unicode 字符缩减到仅 256 个基本符号,彻底消除了 OOV 问题。任何文本都能用这 256 个字节表示。
初始化步骤:
-
对原始文本进行单词切分(如按空格或使用正则分词器)。
-
统计每个单词的频率。
-
对于每个单词,将其字符串编码为 UTF-8 字节序列。Python 中可以通过
word.encode('utf-8')得到bytes对象。 -
将每个字节转换为一个可显示的符号。常见的表示方法有:
- 直接将字节值(0-255)作为符号 ID,但为了可读性,通常映射为字符形式,如
'!'对应 33,而不可打印字节则转为形如<0x00>的字符串或直接保留为bytes对象。 - 在 GPT-2 的实现中,字节被映射到 Unicode 字符(例如 0 -> 'Ā', 1 -> 'ā', ...),但逻辑本质不变。
-
为了简单,我们可以使用整数字节值(0-255)或使用
bytes单字节切片作为符号,但字典键需可哈希。常用做法是将字节转换为整数,但在符号序列中使用整数,便于合并。 -
在每个单词字节序列末尾添加一个特殊的结束符,可以是一个特殊字节值(如 256,或者
</w>的某种表示)。为了不与 0-255 冲突,通常使用超出范围的 ID(如 256 表示结束符)。或者沿用字符串表示,用'</w>'字符串符号,但在字节层面,所有符号都是整数,可定义end_of_word = 256。 -
统计每个字节(0-255 外加结束符)的频率,形成初始词汇表。
Python 实现:
from collections import Counter
def byte_level_initialize(text: str) -> tuple:
"""
对文本进行 Byte-Level BPE 初始化。
返回:
word_symbols: [(byte_sequence, freq), ...] 每个序列为整数列表,以 256 结尾
byte_freq: 各字节的频率字典,键为整数 0-256
"""
# 1. 单词切分(保留标点等,这里简单按空格切分)
words = text.split()
word_freq = Counter(words)
word_symbols = []
byte_freq = Counter()
end_id = 256 # 结束符 ID
for word, freq in word_freq.items():
# 将单词编码为 UTF-8 字节,转为整数列表
byte_seq = list(word.encode('utf-8'))
# 添加结束符
byte_seq.append(end_id)
word_symbols.append((byte_seq, freq))
for b in byte_seq:
byte_freq[b] += freq
return word_symbols, byte_freq
# 示例:
text = "Hello world! Hello"
word_syms, freq = byte_level_initialize(text)
# "Hello" -> [72, 101, 108, 108, 111, 256]
# "world!" -> [119, 111, 114, 108, 100, 33, 256]
# 频率:72:1, 101:1, 108:2, 111:2, 256:2, 等等
后续的 bigram 计算、合并等操作完全类似,只是符号都是整数 ID。新合并的符号可以赋予新的整数 ID(例如从 257 开始递增)。合并时将两个整数 ID a 和 b 替换为新的 ID。最终词汇表就是一系列整数及其对应的字节序列。
优点:
-
完全无 OOV。
-
可以处理任何语言、表情符号、特殊符号。
-
初始符号集极小且固定(256 + 结束符),合并次数可控,最终词汇表大小可任意指定。
注意:在实际解码时,需要能够将整数 ID 序列转换回原始字节,再解码为 UTF-8 字符串。训练过程中需要保存每个 ID(包括合并产生的新 ID)到字节序列的映射。最终解码:将 ID 序列转成字节串,移除结束符对应的字节,然后将字节串 UTF-8 解码为字符串。
在 BPE 训练中,当多个 bigram 频率相同时,如何打破平局?实现一种稳定的选择策略(如基于符号顺序),并体现在代码中¶
BPE 训练需要确定性的行为,即相同输入总是产生相同的合并顺序。当多个 bigram 具有相同的最高频率时,必须有一致的平局打破策略。常用的策略有:
-
按符号对的字典序选择:比较两个符号的字符串表示(或 ID),选择较小的元组。例如
('a', 'b')优先于('a', 'c');若第一个相同,比较第二个。这是最直观且常见的方法,前文第 4 问已采用。 -
按符号对的频率历史或其他启发式规则:比如优先选择最先达到该频率的 pair,但实现复杂。
-
按照训练开始时符号的顺序:如对于字节级,可按 ID 大小排序。
在代码中体现稳定性,关键是当从 bigram_freq 中选出最佳 pair 时,使用确定性的规则。前文实现中,我们通过找到最大频率,然后对具有该频率的所有 pair 取最小字典序,即:
def select_best_pair(bigram_freq: dict) -> tuple:
if not bigram_freq:
return None
max_freq = max(bigram_freq.values())
# 从具有最大频率的 pair 中选择字典序最小的
best_pair = min(pair for pair, f in bigram_freq.items() if f == max_freq)
return best_pair
在 Python 中,元组默认比较就是字典序,因此 min 会优先比较第一个元素,若相同则比较第二个。这保证了无论字典遍历顺序如何,结果唯一。
如果符号不是字符串而是整数 ID(字节级 BPE),字典序就变成了按整数大小排序。这同样稳定。例如 (72, 101) 和 (72, 108) 中前者更小(101 < 108)。
为了更灵活,我们可以定义一个函数,接受 bigram_freq 和可选的排序键:
def select_best_pair_stable(bigram_freq: dict, tie_breaker='lexicographic') -> tuple:
if not bigram_freq:
return None
max_freq = max(bigram_freq.values())
candidates = [pair for pair, f in bigram_freq.items() if f == max_freq]
if tie_breaker == 'lexicographic':
best_pair = min(candidates)
elif tie_breaker == 'reverse_lexicographic':
best_pair = max(candidates)
elif tie_breaker == 'first_seen':
# 如果需要维持插入顺序,可使用 OrderedDict 维护插入次序
# 这里从 bigram_freq 按插入顺序遍历找到第一个最大频率
for pair, f in bigram_freq.items():
if f == max_freq:
best_pair = pair
break
else:
best_pair = min(candidates) # 默认字典序
return best_pair
在 BPE 伪代码中直接调用该稳定选择函数即可。
给定大量文本,为加速 bigram 统计,设计一个高效的数据结构(如字典+计数器),写出更新 bigram 频率的优化步骤¶
当语料非常大时,每次合并后重新扫描所有单词符号序列来计算 bigram 频率会极慢。优化思路是维护一个全局的 bigram 频率字典,并在每次合并时只更新受影响的 bigram 计数,而不是全量重算。
高效数据结构:
-
每个单词的当前符号序列可以存储为一个双向链表或动态数组,但为了便于更新 bigram,通常使用简单的列表,但需要能够快速找到哪些单词包含被合并的 pair。我们可以建立一个映射:
pair -> 包含该 pair 的单词列表(或单词索引),并记录该 pair 在每个单词中的位置(或直接重新扫描该单词)。然而,维护位置会非常复杂。 -
更实用的方法:在每次合并操作后,只修改那些实际发生合并的单词序列,并针对这些序列局部重新计算 bigram 变化。因为一次合并仅影响少数单词(并非所有单词都包含该最高频 pair)。我们可以:
- 维护
word_symbols列表,每个单词的符号序列(list)。 - 维护一个全局
bigram_freq字典,初始通过全量扫描建立。 - 当选定
best_pair = (a, b)后,找出所有包含该 pair 的单词序列(可以在计算频率时同时记录单词索引,例如构建一个pair_to_words映射:pair -> [ (word_idx, count) ],但计数已经是频率加权后的总和,要精确更新需要知道具体每个单词中出现了几次,并考虑单词频率)。 - 对于每个受影响的单词(根据
pair_to_words[best_pair]),在它的序列中执行合并,合并前后需要:- 移除旧序列中所有的 bigram,减去相应频率。
- 添加新序列中产生的 bigram,加上相应频率。
由于单词频率
freq可能大于 1,处理时需要乘以freq。
详细优化步骤:
假设我们有数据结构:
-
word_symbols: 列表,每个元素为[symbols, freq],symbols为可变列表。 -
bigram_freq:defaultdict(int),记录当前全局 bigram 频率。 -
pair_to_word_indices:defaultdict(set),记录每个 pair 出现在哪些单词索引中。这里不考虑单词内部多次出现,仅记录单词索引,对于更新我们可能需要重新扫描该单词,但只扫描被合并的单词,而不是全部。由于一个单词内可能有多个(a,b)出现,重新扫描该单词所有 bigram 更新频率是合理且简单的。
更新步骤(每次合并后):
-
选出
best_pair = (a, b),新符号c = a + b(或新 ID)。 -
获取受影响的单词索引集合:
affected_indices = pair_to_word_indices.get(best_pair, set()),如果为空则忽略(理论上不会)。 -
对于每个索引
idx在affected_indices中: - 取出
symbols, freq = word_symbols[idx]。 - 减去旧 bigram 频率:扫描旧
symbols的所有相邻对,将其频率从bigram_freq中减去freq。同时从pair_to_word_indices[pair]中移除该索引(如果该单词不再包含该 pair,则移除;由于我们后面会重新添加新 pair,可以一次性清除:删除所有旧 pair 的关联,之后重新建立)。 - 应用合并:在
symbols上执行(a,b) -> c的合并(同之前的merge_vocab逻辑),得到new_symbols。 - 添加新 bigram 频率:扫描
new_symbols的所有相邻对,将其频率加上freq。同时将该索引idx加入到对应pair_to_word_indices[pair]集合中。 -
更新
word_symbols[idx]的序列为new_symbols。 -
将
best_pair的频率清零并从pair_to_word_indices中删除,因为所有出现都被合并了(除非合并后由于某种原因重新形成,但在这个局部更新中,旧 pair 已经被消除,新序列中不会包含(a,b)除非c再参与其他合并,但(a,b)作为原始 pair 不会再出现,因为a和b相邻的情况被替换了。注意:如果序列为... a a b ...,合并时扫描从左到右只会替换第一个匹配项?需要确保扫描策略与全量重建一致:一次扫描替换所有非重叠的(a,b)。那么合并后可能仍有剩余的a和b相邻吗?例如a a b:索引 0,1 是a, a不符合;1,2 是a, b符合,替换为a c,最终序列['a', 'c'],不会留下(a,b)。但如果序列是a b a b,合并:c c,没有(a,b)。所以所有(a,b)都会被消除。因此从bigram_freq和映射中删除best_pair是安全的。 -
(可选)如果新产生的 bigram 在
bigram_freq中不存在,则创建;如果某个 bigram 频率降为 0,则从bigram_freq和pair_to_word_indices中删除。
代码示例(局部更新核心逻辑):
def update_bigram_stats(word_symbols, bigram_freq, pair_to_word_indices, best_pair, new_symbol):
a, b = best_pair
# 获取所有包含 best_pair 的单词索引(使用集合避免重复处理)
affected = pair_to_word_indices.get(best_pair, set()).copy()
for idx in affected:
symbols, freq = word_symbols[idx]
# 1. 减去旧 bigram 频率并清除索引关联
for i in range(len(symbols) - 1):
pair = (symbols[i], symbols[i+1])
bigram_freq[pair] -= freq
if bigram_freq[pair] <= 0:
del bigram_freq[pair]
# 从 pair_to_word_indices 中移除 idx(我们稍后重建,简单直接全部清除)
if pair in pair_to_word_indices:
pair_to_word_indices[pair].discard(idx)
if not pair_to_word_indices[pair]:
del pair_to_word_indices[pair]
# 2. 执行合并
new_syms = []
i = 0
while i < len(symbols):
if i < len(symbols) - 1 and symbols[i] == a and symbols[i+1] == b:
new_syms.append(new_symbol)
i += 2
else:
new_syms.append(symbols[i])
i += 1
word_symbols[idx][0] = new_syms # 更新序列
# 3. 添加新 bigram 频率并建立索引
for i in range(len(new_syms) - 1):
pair = (new_syms[i], new_syms[i+1])
bigram_freq[pair] += freq
if pair not in pair_to_word_indices:
pair_to_word_indices[pair] = set()
pair_to_word_indices[pair].add(idx)
# 清除 best_pair 的索引(已被消除)
if best_pair in pair_to_word_indices:
del pair_to_word_indices[best_pair]
这样就可以在 O(受影响单词数 × 平均单词长度) 的时间内完成 bigram 更新,远快于全量扫描。
实现使用优先队列(最大堆)来动态获取最高频 bigram 的 BPE 训练方法,写出合并后如何更新堆中相关 bigram 的逻辑¶
为了更快地获取最高频 bigram,可以使用最大堆。但 Python 的 heapq 是最小堆,可以存储负频率来实现最大堆。堆中的元素为 (-freq, pair),由于堆只保证顶部元素最小(负频率最大),当合并操作改变了某些 bigram 的频率后,堆中对应条目就过时了。有两种处理方式:
-
懒删除(Lazy deletion):不更新堆中元素,而是标记旧频率为无效,或当弹出时检查堆顶频率是否与当前
bigram_freq中记录的一致,不一致则丢弃并重新弹出。这种方法实现简单且常用。 -
主动更新:在堆中直接修改元素,但需要额外索引,复杂且不必要。
懒删除策略:
-
维护
bigram_freq字典作为实时频率的权威源。 -
初始化:将所有
(pair, freq)以(-freq, pair)形式压入列表,然后heapify。 -
每次需要选择最佳 pair 时:
- 从堆中弹出堆顶元素,检查其频率(取负)是否等于
bigram_freq[pair]的当前值。 - 如果相等,说明是有效的最高频 pair;否则丢弃,继续弹出,直到找到有效的。
-
需要注意,堆可能为空(所有 pair 频率为 0 或无 bigram)。
-
选定
best_pair后,执行合并,并使用上一问的局部更新方式更新bigram_freq。 -
更新过程中,某些 pair 的频率增加或减少,甚至新增 pair。对于这些变化的 pair,我们将 新的
(-new_freq, pair)推入堆中,不删除旧条目。旧条目会在未来被弹出时因频率不匹配而丢弃。 -
为了避免堆无限增长,可以定期清理,但通常合并次数有限,堆大小可控。
堆中元素格式:由于可能出现频率相同的情况,还需要实现稳定的平局打破。如果仅存储 (-freq, pair),当频率相同时,Python 会对比 pair 元组来进行堆排序,这正好符合字典序打破平局的要求(因为堆排序基于元组比较,比较完 -freq 后会比较 pair)。但需注意,我们想要的是 最大频率优先,相同频率选最小字典序 pair。由于我们存储的是 (-freq, pair),-freq 越小(即频率越大)排在堆顶;当 -freq 相同时,会按照 pair 排序(默认升序),因此 pair 字典序最小的会在堆顶。这正是我们想要的。例如频率都为 5 的 ('a','b') 和 ('a','c'),对应堆元组 (-5, ('a','b')) 和 (-5, ('a','c')),最小的是 (-5, ('a','b')),正确选出。
实现代码:
import heapq
def bpe_train_with_heap(word_symbols, num_merges):
# 初始化 bigram 频率和堆
bigram_freq = Counter()
pair_to_indices = {}
# 全量扫描建立 bigram_freq 和 pair_to_indices(略,可参考前文)
# ...
# 建立堆
heap = []
for pair, freq in bigram_freq.items():
heapq.heappush(heap, (-freq, pair))
merges = []
for i in range(num_merges):
# 获取有效的最高频 pair
best_pair = None
while heap:
neg_freq, pair = heapq.heappop(heap)
if bigram_freq.get(pair, 0) == -neg_freq:
best_pair = pair
break
# 否则是过时条目,丢弃
if best_pair is None:
break # 没有可合并的 pair
# 合并
a, b = best_pair
new_symbol = a + b if isinstance(a, str) else max_vocab_id + 1 # 简化
merges.append((best_pair, new_symbol))
# 更新 bigram 频率并获取受影响的 pair 及其新频率
affected_pairs = set() # 频率发生变化的所有 pair(旧删除,新添加)
# 执行更新统计(使用上问的 update_bigram_stats 修改版,收集变化的 pair)
# 此处简化为:先获取所有受影响的旧 pair,执行合并后得到新 pairs
# 更新 bigram_freq
update_bigram_stats_with_collect(word_symbols, bigram_freq, pair_to_indices,
best_pair, new_symbol, affected_pairs)
# 将受影响的 pair 的新频率推入堆
for pair in affected_pairs:
new_freq = bigram_freq.get(pair, 0)
if new_freq > 0:
heapq.heappush(heap, (-new_freq, pair))
# 注意:best_pair 的频率变为 0 会被移除,不需推入
return merges
其中 update_bigram_stats_with_collect 需要在修改频率的同时,将旧 pair 和新 pair 均加入 affected_pairs 集合。核心就是前面所述减去旧 bigram 频率,合并,添加新 bigram 频率,并将这些 pair 标记为 affected。由于我们记录了哪些 pair 的频率变动,将其新的频率推入堆中即可。这样堆中会存在同一 pair 的多个条目,但依靠懒删除保证正确性。
给定一个已训练好的 BPE 词汇表(合并规则),实现从词汇表生成合并优先级列表(merges),以便用于编码¶
在实际中,我们可能得到一个词汇表文件,里面包含所有子词 token 及其 ID,但没有显式记录合并的顺序(合并规则)。然而,要正确编码一个单词,我们需要知道合并规则的先后顺序。如果只有最终的词汇表,能否推断出合并顺序?理论上,如果我们知道所有合并产生的 token 以及基础字符,可以通过重建合并树来恢复顺序。但通常 BPE 训练输出都会直接保存 merges 列表(即 (pair, new_token) 的序列),因为它决定了编码行为。
如果手头只有一个词汇表(所有子词),而没有顺序,则无法唯一确定合并顺序,因为不同的合并顺序可能产生相同的词汇表。因此,“从词汇表生成合并优先级列表”这一任务一般是在训练过程中直接记录 merges 并保留。但有些库(如 HuggingFace Tokenizers)会保存一个 merges.txt 文件,其每一行是一对 token,按合并顺序排序(最早合并的在上面)。如果我们只有词汇表,但缺失合并顺序,我们无法恢复。所以这个问题的实际意图是:在实现 BPE 训练器时,需要输出一个合并规则列表,供编码使用。 通常我们训练完后就会拥有 merges,直接存储即可,无需生成。
不过,有一种可能:给定由 BPE 训练过程产生的最终符号集合以及每个符号的组成(例如知道 'low' 是由 'lo' + 'w' 合并而来),我们可以尝试重建合并树,但要确定先后顺序,需要额外的频率信息或时间戳。因此标准做法是:训练时直接记录合并顺序。
代码体现:训练时,在每次合并操作后,将合并的 pair 和新符号追加到 merges 列表中。这就是合并优先级列表,可直接用于编码。例如:
merges = [] # 记录合并顺序
for step in range(num_merges):
best_pair = select_best_pair(...)
new_tok = best_pair[0] + best_pair[1]
merges.append( (best_pair, new_tok) )
# 执行合并...
编码时,按 merges 的顺序逐条应用。这就是之前 encode_word 函数的输入。
实现 BPE 的分词器训练类,包含 train, encode, decode 方法,并支持设置词汇表大小上限¶
我们将实现一个完整的 BPETokenizer 类,整合上述所有内容,采用 Byte-Level 以彻底解决 OOV,并使用堆加速。支持设置词汇表大小上限(vocab_size),初始词汇表大小为 256(字节)+ 1(结束符)= 257,经过 vocab_size - 257 次合并后停止。
类设计:
-
init(self, vocab_size=30000):设置目标词汇表大小。 -
train(self, text):训练 BPE 分词器,生成merges和vocab(id 到字节序列的映射)。 -
encode(self, text):将文本编码为 token ID 列表。 -
decode(self, ids):将 token ID 列表解码为字符串。
内部使用字节级初始化,合并时创建新 ID(从 258 开始)。为了能在 token ID 和字节序列间转换,需要维护一个 id2bytes 字典,记录每个 ID 对应的原始字节序列(便于解码)。合并时,new_id 对应的字节序列是 a_id 的字节序列 + b_id 的字节序列拼接。
详细实现:
import heapq
from collections import Counter, defaultdict
class ByteLevelBPETokenizer:
def __init__(self, vocab_size=30000):
self.vocab_size = vocab_size
self.end_id = 256 # 结束符 ID
self.merges = [] # 合并规则:[(a_id, b_id, new_id), ...]
self.id2bytes = {} # ID -> bytes
self.unk_id = None # Byte-level 不需要 UNK,但可保留
def train(self, text: str):
# 1. 单词分词与频率统计
words = text.split()
word_freq = Counter(words)
# 2. 初始化字节序列
word_symbols = [] # [[int_list], freq]
bigram_freq = Counter()
pair_to_indices = defaultdict(set)
# 初始化 id2bytes: 0-255 为单字节
for i in range(256):
self.id2bytes[i] = bytes([i])
# 结束符用空字节表示,但为了方便解码,可以赋予一个特殊字节,这里我们不在 id2bytes 中存放结束符,解码时遇到结束符就忽略
# 也可以将结束符映射到一个不会出现的字节,但最终解码我们会处理。
for word, freq in word_freq.items():
byte_seq = list(word.encode('utf-8')) + [self.end_id]
word_symbols.append([byte_seq, freq])
for i in range(len(byte_seq) - 1):
pair = (byte_seq[i], byte_seq[i+1])
bigram_freq[pair] += freq
pair_to_indices[pair].add(len(word_symbols) - 1)
# 3. 初始化堆
heap = []
for pair, f in bigram_freq.items():
heapq.heappush(heap, (-f, pair))
# 当前最大 ID
next_id = 257 # 0-256 保留
# 目标合并次数
num_merges = self.vocab_size - next_id # 因为 257 是当前词汇量
if num_merges < 0:
num_merges = 0
for step in range(num_merges):
# 弹出有效最高频 pair
best_pair = None
while heap:
neg_f, pair = heapq.heappop(heap)
if bigram_freq.get(pair, 0) == -neg_f:
best_pair = pair
break
if best_pair is None:
break
a, b = best_pair
new_id = next_id
next_id += 1
# 记录合并规则
self.merges.append((a, b, new_id))
# 新 ID 的字节序列 = a 的字节序列 + b 的字节序列
self.id2bytes[new_id] = self.id2bytes[a] + self.id2bytes[b]
# 更新受影响单词的 bigram 统计
affected = pair_to_indices.get(best_pair, set()).copy()
changed_pairs = set()
for idx in affected:
symbols, freq = word_symbols[idx]
# 减去旧 bigram
for i in range(len(symbols) - 1):
p = (symbols[i], symbols[i+1])
bigram_freq[p] -= freq
if bigram_freq[p] <= 0:
del bigram_freq[p]
if p in pair_to_indices:
pair_to_indices[p].discard(idx)
if not pair_to_indices[p]:
del pair_to_indices[p]
changed_pairs.add(p)
# 合并
new_syms = []
i = 0
while i < len(symbols):
if i < len(symbols) - 1 and symbols[i] == a and symbols[i+1] == b:
new_syms.append(new_id)
i += 2
else:
new_syms.append(symbols[i])
i += 1
word_symbols[idx][0] = new_syms
# 添加新 bigram
for i in range(len(new_syms) - 1):
p = (new_syms[i], new_syms[i+1])
bigram_freq[p] += freq
if p not in pair_to_indices:
pair_to_indices[p] = set()
pair_to_indices[p].add(idx)
changed_pairs.add(p)
# 清理 best_pair 的索引
if best_pair in pair_to_indices:
del pair_to_indices[best_pair]
changed_pairs.discard(best_pair)
# 将改变的 pair 的新频率推入堆
for p in changed_pairs:
f = bigram_freq.get(p, 0)
if f > 0:
heapq.heappush(heap, (-f, p))
# 最终词汇表大小
self.vocab = self.id2bytes.copy() # 完整的 id -> bytes 映射
def encode(self, text: str) -> list:
# 按单词处理,每个单词编码后拼接,插入结束符 ID
words = text.split()
token_ids = []
for word in words:
seq = list(word.encode('utf-8')) + [self.end_id]
for a, b, new_id in self.merges:
new_seq = []
i = 0
while i < len(seq):
if i < len(seq)-1 and seq[i] == a and seq[i+1] == b:
new_seq.append(new_id)
i += 2
else:
new_seq.append(seq[i])
i += 1
seq = new_seq
token_ids.extend(seq)
return token_ids
def decode(self, ids: list) -> str:
# 将所有 ID 转成字节序列,拼接,然后处理结束符
byte_list = []
for tid in ids:
if tid == self.end_id:
# 词尾,添加空格
byte_list.append(b' ')
else:
byte_list.append(self.id2bytes[tid])
full_bytes = b''.join(byte_list)
# 解码为 UTF-8,去除末尾多余空格
text = full_bytes.decode('utf-8', errors='replace')
# 可能结尾有空格(最后一个词尾),strip 去掉首尾,但保留中间空格
text = ' '.join(text.split())
return text
测试:
tokenizer = ByteLevelBPETokenizer(vocab_size=300)
tokenizer.train("low low low lower lowest")
tokens = tokenizer.encode("lowest")
print(tokens)
print(tokenizer.decode(tokens))
该类完整实现了训练、编码、解码,并支持词汇表大小上限,采用 Byte-Level 彻底解决 OOV。
写出在训练 BPE 时,如何正确处理数字和标点:是否分词为独立符号?设计预处理步骤并融入训练伪代码¶
数字和标点的处理直接影响最终子词单元的质量。目标是将文本分割成合适的初始单词,使得标点、数字等能够作为独立的实体参与 BPE 合并,而不是粘在相邻单词上,导致词汇表膨胀或产生无意义的子词。常见预处理策略有:
-
按空格分词,但将标点与相邻单词分开:在英文 NLP 中,通常使用 Moses tokenizer 或类似方法,将标点(如
.,!?;:()[]{}'"等)切分为独立的 token,数字保留原样或进一步切分(如将年份1990s分为1990和s)。这可以确保标点成为独立的“单词”,并加上结束符,使得 BPE 能够学习到标点后通常跟随词尾等特点。 -
对于数字,可以将连续数字作为一个整体单词,也可以按位拆分。但通常保留连续数字作为整体,让 BPE 内部按需合并(如
100可能合并为100或者整体)。如果语料中有很多数字,BPE 可能学习到100的整体表示,但更好的做法是让数字保持为字符级或字节级,由 BPE 自动处理。 -
预处理伪代码:在训练前加入一个分词器(tokenizer),将原始文本拆分成单词序列(每个单词将独立加上
</w>)。这个分词器的行为是:按空格和标点拆分,保留标点作为单独 token。
具体预处理步骤:
-
利用正则表达式将文本分割为 tokens(单词和标点独立)。例如,模式
r'\w+|[^\w\s]'可以匹配所有单词字符序列或者非空白非单词字符(标点)。也可以处理缩写中的标点如's等。 -
对于每个 token,统计频率,将它们作为“单词”输入 BPE 初始化。即 token 序列中的每一项都看作一个单词,末尾加结束符。
-
对于数字,可以选择保留连续数字,或者进一步拆分为单个数字。但 BPE 在字符级或字节级初始化后,自然会处理数字内部的合并。若初始化为字符级,数字字符
'1','2'等成为基本符号;若为 Byte-Level,数字的 UTF-8 编码为一个字节(ASCII),仍是基本符号。
融入训练伪代码:
算法 BPE_Train_with_preprocess(raw_text, num_merges):
# 预处理:tokenize 成单词和标点
tokens = []
for match in re.finditer(r'\w+|[^\w\s]', raw_text):
tokens.append(match.group())
# 或者更复杂的如 Moses tokenizer
# 统计频率
word_freq = Counter(tokens)
# 后续初始化与合并同原始伪代码
word_symbols = ...
for word, freq in word_freq.items():
symbols = list(word) + ['</w>']
...
在 Byte-Level BPE 中,预处理同样重要,因为需要将字节序列分配给正确的“单词”边界。如果文本没有预先将标点分离,"world!" 会被视为一个单词,结束符加在感叹号后。经过 BPE 可能学到 world!</w> 作为一个整体,不利于泛化。分离后 "world" 和 "!" 各有词尾,BPE 会学到 world</w> 和 !</w>,更为合理。
数字处理细节:若数字保留为连续 token(如 "123"),BPE 可能将其整体作为一个符号或拆分为多字节合并。由于数字种类无限,通常 BPE 会将连续数字拆成单数字字符,依赖模型学习数字的组成模式。因此,一些实现会将数字预处理为 0 或保留原样。但如果保持 Byte-Level,字节序列 49 50 51(即 '1','2','3')会自然学习合并。
标点处理总结:将标点与单词分开作为独立 token 是标准做法。在训练伪代码中,第一步就是应用一个合适的 tokenizer 将文本划分为 tokens,再用 BPE 训练这些 tokens。这样可以避免标点粘附在词上导致的词汇冗余。
实现一个函数,输入已训练的 BPE 模型和一段文本,输出该文本中每个 token 的 ID 序列(需包含词汇表到 ID 的映射)¶
此函数的目标是将文本转化为 token ID 列表,以便喂给下游模型。为此,需要一个完整的词汇表到 ID 的映射(vocab_to_id)。该映射在训练后确定,包含所有最终符号(初始字符/字节、结束符、合并产生的新符号)。对于 Byte-Level BPE,符号是整数 ID,但通常模型输入需要连续的 token ID(从 0 到 vocab_size-1)。我们需要输出从子词符号到 ID 的映射,以及编码得到的 ID 序列。
实现思路:
-
训练完成后,我们拥有所有合法符号及其对应的 ID(
symbol_to_id)。对于字符级 BPE,符号是字符串;对于 Byte-Level BPE,符号最终是整数 ID(代表字节或合并的 token)。为了统一,我们构建一个字典,将每个 token 的字符串表示(或整数 ID)映射到一个整数索引。 -
编码函数
encode_to_ids首先调用已有的encode得到 token 序列(符号列表),然后通过映射表转为 ID 序列。如果使用的是 Byte-Level BPE,encode返回的可能已经是整数 ID,但那些 ID 是合并过程中分配的原始值(从 0 开始的字节、结束符 256、新 ID 257+)。我们可以选择将这些原始 ID 直接作为模型输入的 token ID(此时映射表就是原始 ID 到自身的映射),但更常见的是重新映射到连续的 0..V-1。不管哪种,都需要生成映射表。 -
此函数要求“需包含词汇表到 ID 的映射”,因此我们应该返回 token ID 列表以及映射字典。
示例实现(基于字符级 BPE 和词汇表):
def encode_with_ids(encoded_tokens: list, token_to_id: dict) -> tuple:
"""
将 token 序列转换为 ID 序列,同时返回映射表。
参数:
encoded_tokens: 编码后的 token 列表(字符串)
token_to_id: {token: id} 映射字典
返回:
ids: 整数 ID 列表
vocab_to_id: 词汇表到 ID 的映射(与输入相同,便于调用者获取)
"""
ids = [token_to_id[token] for token in encoded_tokens]
return ids, token_to_id
如果尚未建立映射,可在 BPE 训练结束时构建:
def build_vocab_map(word_symbols, merges, end_token='</w>'):
"""构建最终词汇表到 ID 的映射"""
vocab = set()
for syms, _ in word_symbols:
for s in syms:
vocab.add(s)
# 也加入所有合并产生的新符号(已在 word_symbols 中体现,但以防万一)
for (a, b), new_sym in merges:
vocab.add(new_sym)
# 按某种顺序排序,例如按长度然后字典序
sorted_vocab = sorted(list(vocab), key=lambda x: (len(x) if isinstance(x, str) else 0, x))
token_to_id = {tok: i for i, tok in enumerate(sorted_vocab)}
return token_to_id
对于 Byte-Level BPE,符号本身就是整数 ID,可将其直接作为模型输入,但需确保 ID 连续。通常原始合并中 ID 从 0 到 vocab_size-1 连续,可以不加转换。此时 token_to_id 就是恒等映射 {i: i for i in range(vocab_size)}。
完整函数示例:
def encode_text_to_ids(text: str, tokenizer) -> tuple:
"""
使用训练好的 tokenizer 编码文本,返回 ID 序列和词汇表映射。
tokenizer 需提供 encode 方法返回 token 列表,以及 vocab_to_id 属性。
"""
tokens = tokenizer.encode(text) # 返回 token 列表(字符串或整数)
ids = [tokenizer.token_to_id[tok] for tok in tokens]
return ids, tokenizer.token_to_id
注意:特殊 token(如 [CLS], [SEP])也需要在映射表中,可在训练后添加。
给定一组语料,实现添加特殊 token(如 [CLS], [SEP], [MASK], [PAD], [UNK])到词汇表中,并确保它们在 BPE 训练过程中不被拆分¶
特殊 token 是完整的、不可分割的符号。在 BPE 训练中,若不特殊处理,[CLS] 会被初始化为字符序列 ['[', 'C', 'L', 'S', ']'],然后可能参与合并,甚至被拆散,这是不可接受的。因此需要保证这些特殊标记在训练时就作为单个整体 token,并且永远不被进一步拆分或合并。通常采用以下方法:
-
在初始化前,将特殊 token 作为独立“单词”加入词频统计:例如,在语料中可能没有
[CLS],但我们可以人为地将其加入词频字典,频率可以设为 0 或 1(如果语料中本身出现了[CLS]文本则正常统计)。 -
将整个特殊 token 视为一个不可拆分的符号:即它由单一符号表示,而不是字符序列。初始化时,我们不将其拆分为字符加结束符,而是直接作为一个符号,末尾也加上结束符
</w>,使其保持完整性。例如,单词"[CLS]"被表示成符号序列['[CLS]', '</w>']。这样在 BPE 合并过程中,[CLS]作为一个原子符号存在,只有它与</w>的边界可能参与合并,但因为[CLS]本身是整体,不会被拆分。 -
如果需要完全阻止它与任何其他符号合并(包括结束符),可以将
[CLS]视为已包含词尾的特殊 token,即不需要额外的</w>,或者在合并算法中跳过包含特殊 token 的 pair。通常的做法是:特殊 token 不需要结束符,它们自身就是边界。可在初始化时把它们当成完整的 token,序列只有['[CLS]']而没有</w>,但这样一来它们永远无法与后序单词合并(符合预期)。另外,合并规则中如果涉及特殊 token,应禁止合并。实现时可以在选择 best pair 时排除包含特殊 token 的 pair,或者在合并步骤中不对它们进行替换。最简单的策略是:在初始化时,特殊 token 被封装为单独的符号,如'[CLS]',并且该符号不在任何单词序列中作为与其他符号可合并的单元(因为它的序列就是['[CLS]'],不与任何其他符号相邻)。但为了编码文本时灵活插入[CLS],我们可以定义[CLS]就是一个 token,编码时如果用户要求,直接加入输出,而不经过字符拆分与合并。
实现步骤:
-
定义特殊 token 列表,如
special_tokens = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]"]。 -
在词频统计阶段,为每个特殊 token 分配一个极小的频率(或 0),使得它们出现在词汇表中。或者可以在初始化词汇表时直接插入它们。
-
初始化符号序列时,对于特殊 token 单词,不进行字符拆分,而是直接将其整个字符串作为一个符号,并加上
</w>(或直接用自身作为词尾,免去结束符)。例如,"[CLS]"->['[CLS]', '</w>']。 -
在 BPE 合并循环中,我们需要确保合并不会破坏特殊 token 的内部。如果我们将特殊 token 作为原子符号,它就不会被拆分。但必须防止合并操作将特殊 token 与相邻符号合并。例如
'[CLS]'和'</w>'可能会被认为是一个 pair。如果我们保留结束符,那么 BPE 可能会学习到合并('[CLS]', '</w>')成为一个新符号"[CLS]</w>"。这并无大碍,但通常我们希望[CLS]作为独立的 token 出现在最终词汇表中。为了避免合并,我们可以选择在每次选择 best pair 时,排除任何一个元素在special_tokens集合中的 pair。或者在合并后,如果生成了包含特殊 token 的新符号,禁止它。更简单的是:在初始化时,特殊 token 不带结束符,且序列只包含该 token 本身(无结束符),这样它就处于孤立状态,不会与任何符号形成 bigram(除非单词内有多个 token?但我们不会把特殊 token 和普通单词连在一起)。在编码时,特殊 token 直接映射到一个 ID,不需要经过 BPE 合并流程。
示例实现:
def initialize_vocab_with_special(word_freq: dict, special_tokens: list) -> list:
word_symbols = []
special_token_set = set(special_tokens)
for word, freq in word_freq.items():
if word in special_token_set:
# 特殊 token 作为一个整体符号,不加结束符,避免与其他合并
symbols = [word]
else:
symbols = list(word) + ['</w>']
word_symbols.append((symbols, freq))
# 确保所有特殊 token 都存在于 word_symbols 中,即使频率为 0
existing = {word for word, _ in word_freq.items()}
for sp in special_tokens:
if sp not in existing:
word_symbols.append(([sp], 0))
return word_symbols
在合并循环中,best_pair 如果包含了特殊 token 的元素,则跳过(但如果我们采用无结束符的方式,特殊 token 根本不与任何符号相邻,故其 bigram 频率为 0,不可能被选中)。如果特殊 token 有结束符,我们可在选择时过滤:
def select_best_pair(bigram_freq, forbidden_symbols):
# 过滤掉包含禁止符号的 pair
allowed = {pair: freq for pair, freq in bigram_freq.items()
if pair[0] not in forbidden_symbols and pair[1] not in forbidden_symbols}
if not allowed:
return None
max_freq = max(allowed.values())
best_pair = min(pair for pair, f in allowed.items() if f == max_freq)
return best_pair
最终:词汇表构建时,特殊 token 将被分配固定的 ID(通常是最前面的几个 ID),且编码时若文本中出现这些特殊 token 字符串,会直接映射到对应 ID,无需字符拆分。
实现从 HuggingFace 格式的 tokenizer.json 或 vocab.json 中加载 BPE 模型,并完成 encode/decode 的基本功能(主要解析合并规则)¶
HuggingFace 的 tokenizers 库中,BPE 分词器的模型文件通常包含 vocab.json 和 merges.txt(或一并打包在 tokenizer.json 中)。vocab.json 记录了 token 到 ID 的映射,merges.txt 记录了合并规则,每行是一对 token(用空格分隔),按顺序排列(最早合并的在第一行)。我们需要解析这些文件并构建编码/解码逻辑。
解析步骤:
-
读取
merges.txt:文件头可能包含版本注释,真正的合并规则从某一行开始。标准格式是#version: 0.2然后换行之后每行token_a token_b。我们将这些行解析为合并列表,每项为(token_a, token_b)。合并顺序即为文件中的出现顺序。 -
读取
vocab.json:得到{token: id}字典。同时,它也包含了所有基础字符和合并产生的子词。注意这个 vocab 的 ID 可能不连续或有特殊 token。我们需要反向映射id_to_token用于解码。 -
构建 token 到 ID 映射:直接用
vocab.json的字典。 -
编码:与之前类似,输入单词,拆分为字符序列,加上
</w>表示结束(HuggingFace 使用</w>作为词尾,也有用Ġ表示空格的变体,需根据实际情况调整)。然后按照 merges 顺序逐条合并。注意 HuggingFace 的 BPE 可能在字符前会有一个特殊空格字符,例如"hello"可能被表示为"h" "e" "l" "l" "o",而词首的空格用Ġ(GPT-2) 表示。但 BPE 基础模型通常使用</w>作为词尾,并在文本开头不加空格。我们需要明确是哪一种实现。常见的是 RoBERTa / GPT-2 使用的 byte-level BPE 与原始 BPE 稍有不同,但merges.txt的解析逻辑相同。此处以标准原始 BPE(带</w>)为例。
实现代码:
import json
def load_hf_bpe(vocab_path, merges_path):
# 加载词汇表
with open(vocab_path, 'r', encoding='utf-8') as f:
vocab = json.load(f) # token -> id
id_to_token = {v: k for k, v in vocab.items()}
# 加载合并规则
merges = []
with open(merges_path, 'r', encoding='utf-8') as f:
for line in f:
line = line.strip()
if not line or line.startswith('#'):
continue
parts = line.split()
if len(parts) == 2:
merges.append(tuple(parts))
return vocab, id_to_token, merges
def encode_with_merges(word, merges, end_suffix='</w>'):
symbols = list(word) + [end_suffix]
for a, b in merges:
new_symbols = []
i = 0
while i < len(symbols):
if i < len(symbols)-1 and symbols[i] == a and symbols[i+1] == b:
new_symbols.append(a+b)
i += 2
else:
new_symbols.append(symbols[i])
i += 1
symbols = new_symbols
return symbols
def encode_text(text, merges, vocab, end_suffix='</w>'):
# 分词,这里简单按空格分,实际需预处理与训练时一致
words = text.split()
token_ids = []
for word in words:
tokens = encode_with_merges(word, merges, end_suffix)
for tok in tokens:
token_ids.append(vocab[tok])
return token_ids
def decode_ids(ids, id_to_token, end_suffix='</w>'):
tokens = [id_to_token[i] for i in ids]
text = ''.join(tokens).replace(end_suffix, ' ').strip()
return text
如果是对 tokenizer.json,该 JSON 文件结构较复杂,包含 model 字段里的 vocab 和 merges(可能是数组)。可以类似解析:
def load_from_tokenizer_json(path):
with open(path, 'r') as f:
data = json.load(f)
model = data['model']
vocab = model['vocab'] # token -> id
merges = model['merges'] # 列表,元素如 "a b"
# 处理 merges 字符串
merges_list = []
for m in merges:
parts = m.split()
if len(parts) == 2:
merges_list.append(tuple(parts))
id_to_token = {v: k for k, v in vocab.items()}
return vocab, id_to_token, merges_list
其余编码解码逻辑相同。
编写一个函数,评估 BPE 分词器的压缩率:计算原文 UTF-8 字节数与分词后 token 序列字节数(按词汇表大小所需比特数估算)之比¶
压缩率可以体现 BPE 对于文本的压缩能力。给定一段文本,我们可以分别计算原始 UTF-8 编码的比特数,以及用分词器分词后,如果对每个 token 使用等长编码(即用固定位数表示词汇表中的每个 token)所需比特数,并计算比率(原始比特数 / 压缩后比特数)。比率大于 1 表示压缩有效。
评估方法:
-
原始比特数:
original_bytes * 8 -
压缩后比特数:若词汇表大小为
V,每个 token 理论上需要ceil(log2(V))比特来唯一表示(固定长度编码)。这是理论下界(实际如 Huffman 编码可更优,但这里用等长编码估算)。分词后得到N个 token,所需比特数 ≈N * ceil(log2(V))。 -
压缩率 = 原始比特数 / 压缩后比特数。
实现细节:
-
使用训练好的分词器对文本进行编码,得到 token 序列(或 ID 序列,数量 N)。
-
获取词汇表大小
V(即vocab_size)。 -
计算
bits_per_token = math.ceil(math.log2(V)),如果 V=1,则 bits_per_token=1。 -
compressed_bits = len(token_ids) * bits_per_token -
原始文本转为 UTF-8 字节:
original_bytes = len(text.encode('utf-8')),比特数为original_bits = original_bytes * 8。 -
返回
original_bits / compressed_bits以及相关信息。
注意事项:
-
词汇表大小通常包含结束符、特殊 token 等,应与编码时的 token 集一致。
-
如果词汇表非常大,
bits_per_token会很高,但 token 数量会减少,权衡体现压缩能力。 -
如果对比不同分词器,应在同一文本上评估。
代码示例:
import math
def evaluate_compression_ratio(text: str, tokenizer) -> dict:
"""
计算 BPE 分词器的压缩率。
返回:{'original_bits': int, 'compressed_bits': int, 'ratio': float, 'num_tokens': int}
"""
# 原始比特数
original_bytes = len(text.encode('utf-8'))
original_bits = original_bytes * 8
# 编码
token_ids = tokenizer.encode(text) # 返回 ID 列表
num_tokens = len(token_ids)
# 词汇表大小(假设 tokenizer.vocab_size 存在)
V = tokenizer.vocab_size
if V <= 1:
V = 2 # 防止 log2(1)=0
bits_per_token = math.ceil(math.log2(V))
compressed_bits = num_tokens * bits_per_token
ratio = original_bits / compressed_bits if compressed_bits > 0 else float('inf')
return {
'original_bits': original_bits,
'compressed_bits': compressed_bits,
'ratio': ratio,
'num_tokens': num_tokens,
'vocab_size': V
}
对于 Byte-Level BPE,vocab_size 可能为 30000,bits_per_token=15,num_tokens 相比原始字节数减少很多,通常压缩率在 2~4 之间,表明 BPE 在减少序列长度的同时保持了词汇可控。
处理大规模语料时内存不足,实现分批 BPE 训练的伪代码:分块读入语料,逐块统计词频和 bigram,最后合并全局 bigram 计数进行合并¶
当语料过大无法全部加载内存时,需要分批处理。BPE 训练的第一步是词频统计,可以在遍历语料时累积。如果我们能分批统计词频并合并,即可得到全局词频。但 BPE 的合并步骤需要基于全局词频和符号序列,如果内存无法容纳所有单词的符号序列(例如单词种类极多),我们需要采用一些近似或外部存储策略。标准的分批训练方法常用于大规模 BPE(如 SentencePiece),使用如下策略:
-
分批统计词频:逐块读取文本,切词,更新词频字典(内存中保留频率即可,单词种类通常远小于总 token 数,可以存入内存)。如果单词种类过多,可以牺牲罕见词进行过滤(如 min frequency)。
-
初始化的符号序列:在获得全局词频后,我们可以为每个单词生成初始符号序列。如果单词种类依然太多,可在合并阶段部分驻留磁盘,或采用高效的流式合并算法。但常见的做法是:一旦词频统计完毕,单词种类数可能为几十万到几百万,符号序列可以用字典存储,对于现代内存尚可接受。如果单词种类大到百亿级别,则需采用分布式处理。对于“内存不足”通常指无法一次性加载文本,但词频字典还是能存的。伪代码可以这样设计。
分批统计词频:
procedure batch_collect_word_freqs(corpus_files):
word_freq = empty dictionary
for each file in corpus_files:
for line in file:
words = tokenize_line(line) # 按空格或预分词
for w in words:
word_freq[w] += 1
return word_freq
之后使用全局 word_freq 进行初始化。
如果连单词的符号序列都放不下(即 word_symbols 太大),我们可以采取 分片合并策略:将单词根据首字母或哈希分桶,对每个桶独立进行部分 BPE 合并?但合并需要跨单词统计 bigram,不能完全独立。真正的超大规模通常使用分布式计算框架(如 Spark)或 SentencePiece 内部基于后缀数组的算法,不是基于单词序列。但对于本问题,我们可在伪代码中体现基本的分批思想:
-
使用外部排序或数据库存储单词符号序列,每次计算 bigram 频率时扫描磁盘上的序列,合并后写回。这属于 I/O 密集型,但可行。
-
伪代码简化:假设词频统计后单词种类可控,其余步骤在内存进行。若不能,可采用“多次遍历”方法:第一遍只统计词频,第二遍初始化符号序列并写入临时文件,第三遍开始合并,每次合并读取临时文件,更新并写回。这可以处理任意大小的词表,虽然速度慢。伪代码描述如下:
procedure incremental_merge(word_freq, num_merges):
# 假设 word_freq 大小可放入内存,但符号序列总字节很大,我们用文件存储
write_initial_word_symbols_to_file(word_freq, "symbols.txt")
for step in 1..num_merges:
bigram_freq = {}
# 读取 symbols.txt,统计 bigram 频率
for each (symbols_str, freq) in read_symbols_file("symbols.txt"):
symbols = parse(symbols_str)
for i in 0..len(symbols)-2:
pair = (symbols[i], symbols[i+1])
bigram_freq[pair] += freq
if bigram_freq empty: break
best_pair = select_best(bigram_freq)
merges.append(best_pair)
# 再次读取 symbols.txt,对每个序列合并 best_pair,写入新临时文件
new_file = "symbols_step_" + step + ".txt"
for each (symbols_str, freq) in read_symbols_file("symbols.txt"):
symbols = merge_in_sequence(parse(symbols_str), best_pair)
write new_file with (symbols, freq)
替换 symbols.txt 为新文件
return merges, final_vocab_from_last_file
这种伪代码展示了如何分批/外存处理。
但鉴于问题可能期待的是“分批统计词频和 bigram”,更可能在合并前,初始 bigram 频率可以通过分批统计得到,而不需要一次性构建所有符号序列。实际上,初始 bigram 频率可以直接从词频字典分批构建:对于每个单词,其字符序列已知,可以流式生成初始 bigram,累加到全局字典。这一步可以分批进行:将单词列表分块,每块生成其初始符号序列并更新 bigram 频率字典,然后丢弃该块的符号序列,只保留 bigram 计数。但后续合并需要修改符号序列,所以不能全部丢弃。因此需要保留符号序列以便后续合并。若内存不足以保留所有符号序列,可能采用前述外存方法。
这里提供一个更务实的“分批 BPE 训练”伪代码,假设内存足够存下符号序列但文本不能一次性加载:
算法 Batched_BPE(corpus_generator, num_merges):
word_freq = Counter()
# 批次处理文本
for text_chunk in corpus_generator:
words = tokenize(text_chunk)
word_freq.update(words)
# 初始化(内存中)
word_symbols = []
for word, freq in word_freq.items():
symbols = list(word) + ['</w>']
word_symbols.append([symbols, freq])
# 常规合并(内存中完成)
return standard_bpe_merge(word_symbols, num_merges)
表明主要瓶颈是文本读取,通过生成器分批读取即可。
在 BPE 训练时,实现一个停止条件:不仅基于合并次数,也可基于词汇表大小达到预设值,修改循环条件¶
停止条件可以由用户指定一个目标词汇表大小 target_vocab_size,而不是硬性的合并次数。当合并产生的词汇数量达到这个值时就停止。初始词汇表大小 V0 是已知的(字符级:字符种类数 + 1(结束符);字节级:256 + 1 = 257)。每次合并会增加一个新符号,因此词汇表大小增加 1。所以要达到目标词汇表大小,需要的合并次数为 target_vocab_size - V0。也可以循环时既检查次数限制,又检查词汇表大小。
实现修改:
def train_bpe(word_freq, target_vocab_size=None, num_merges=None):
# 初始化...
word_symbols, vocab = initialize(word_freq)
V = len(vocab) # 当前词汇表大小
max_merges = float('inf')
if target_vocab_size:
max_merges = target_vocab_size - V
if num_merges is not None:
max_merges = min(max_merges, num_merges)
merges = []
for step in range(max_merges):
bigram_freq = compute_bigram_freq(word_symbols)
if not bigram_freq:
break
best_pair = select_best_pair(bigram_freq)
new_symbol = best_pair[0] + best_pair[1]
merges.append((best_pair, new_symbol))
word_symbols = merge_vocab(word_symbols, best_pair)
vocab.add(new_symbol)
V += 1
# 若已达到目标词汇大小可提前 break,但 max_merges 已保证
return merges, vocab, word_symbols
停止条件多样:也可以基于字符覆盖率达到某个阈值,或最高频 bigram 的频率低于某个阈值(避免罕见合并)。但词汇表大小是最常见的控制参数,直接转换为合并次数即可。
整合核心功能,写出一个简洁的 BPE 分词器完整训练和推理的伪代码,包含训练(含词频统计、bigram 合并)、保存模型、编码、解码的全部流程¶
综合前面的所有要素,下面给出一个完整的 BPE 分词器伪代码,采用 Byte-Level(解决 OOV)并支持特殊 token、词汇表大小上限、保存和加载模型。
完整流程伪代码:
类 BPETokenizer:
def __init__(self, vocab_size=30000, special_tokens=[]):
self.vocab_size = vocab_size
self.special_tokens = special_tokens # e.g., ["[PAD]","[UNK]","[CLS]","[SEP]"]
self.end_id = 256
self.merges = [] # [(a_id, b_id, new_id), ...]
self.id_to_bytes = {} # id -> bytes
self.token_to_id = {} # string/int token -> final model id
self.id_to_token = {}
def train(self, text_iterable):
# 1. 词频统计(分批)
word_freq = Counter()
for text_chunk in text_iterable:
words = text_chunk.split()
word_freq.update(words)
# 2. 加入特殊 token 频率(保证存在)
for sp in self.special_tokens:
if sp not in word_freq:
word_freq[sp] = 0
# 3. Byte-Level 初始化
word_symbols = [] # [(id_seq, freq)]
pair_to_indices = defaultdict(set)
bigram_freq = Counter()
for i in range(256):
self.id_to_bytes[i] = bytes([i])
for word, freq in word_freq.items():
if word in self.special_tokens:
# 特殊 token:分配新 ID,作为一个整体,不加结束符
new_id = len(self.id_to_bytes)
self.id_to_bytes[new_id] = word.encode('utf-8')
seq = [new_id] # 不加结束符
else:
byte_seq = list(word.encode('utf-8')) + [self.end_id]
seq = byte_seq
word_symbols.append([seq, freq])
# 更新 bigram 频率和索引
for i in range(len(seq)-1):
pair = (seq[i], seq[i+1])
bigram_freq[pair] += freq
pair_to_indices[pair].add(len(word_symbols)-1)
# 4. 建立堆
heap = []
for pair, f in bigram_freq.items():
heapq.heappush(heap, (-f, pair))
next_id = len(self.id_to_bytes)
num_merges = self.vocab_size - next_id
if num_merges < 0: num_merges = 0
# 5. 合并循环
for step in range(num_merges):
# 获取有效最佳 pair
best_pair = None
while heap:
neg_f, pair = heapq.heappop(heap)
if bigram_freq.get(pair, 0) == -neg_f:
best_pair = pair
break
if best_pair is None:
break
a, b = best_pair
new_id = next_id
next_id += 1
self.merges.append((a, b, new_id))
self.id_to_bytes[new_id] = self.id_to_bytes[a] + self.id_to_bytes[b]
# 更新受影响单词的 bigram 统计(同前详述,此处简写)
update_affected_words(word_symbols, best_pair, new_id, bigram_freq, pair_to_indices, heap)
# 6. 构建最终 token -> id 映射
# 我们已经有 self.id_to_bytes (id->bytes),可将其映射为 token 字符串
# 为方便使用,构建 token_to_id 字典(以 bytes 的字符串表示或保留原 ID)
# 如果直接使用 ID 作为模型输入,则 token_to_id 为 id -> id
for id_ in range(next_id):
# 将字节解码为字符串(可能包含不可打印字符,用特殊表示)
# 简单做法:token 就是该 ID 对应的字节串,但用作字符串可能含乱码
# 在推理时,我们直接使用 ID 列表,不必转为字符串。
pass
# 最终 vocab_size = next_id
self.vocab_size_actual = next_id
def encode(self, text):
words = text.split()
ids = []
for word in words:
if word in self.special_tokens:
# 直接获取其 ID(假设在训练时已分配)
# 此处在训练时我们为特殊 token 分配了 ID,可以通过查找得到
sp_id = self._get_special_id(word)
ids.append(sp_id)
continue
seq = list(word.encode('utf-8')) + [self.end_id]
for a, b, new_id in self.merges:
new_seq = []
i = 0
while i < len(seq):
if i < len(seq)-1 and seq[i] == a and seq[i+1] == b:
new_seq.append(new_id)
i += 2
else:
new_seq.append(seq[i])
i += 1
seq = new_seq
ids.extend(seq)
return ids
def decode(self, ids):
byte_buf = []
for tid in ids:
if tid == self.end_id:
byte_buf.append(b' ')
else:
byte_buf.append(self.id_to_bytes[tid])
text = b''.join(byte_buf).decode('utf-8', errors='replace')
return ' '.join(text.split())
def save(self, path_prefix):
# 保存 id_to_bytes (pickle or json),merges,special_tokens 等
...
def load(self, path_prefix):
# 加载模型参数
...
伪代码总结:
-
训练:分批读文本 -> 统计词频 -> 字节级初始化(特殊 token 整词处理) -> 堆加速的合并循环直到词汇表达到设定大小 -> 记录合并规则和 ID 映射。
-
编码:单词按字节拆分 -> 应用 merges 得到子词 ID 序列 -> 输出。
-
解码:ID 序列 -> 查找字节拼接 -> 替换结束符为空格 -> UTF-8 解码成文本。
-
保存/加载:序列化
merges,id_to_bytes, 特殊 token 列表等。
此框架整合了前述所有核心功能,构成一个实用的 BPE 分词器。