跳转至

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 个字节表示。

初始化步骤:

  1. 对原始文本进行单词切分(如按空格或使用正则分词器)。

  2. 统计每个单词的频率。

  3. 对于每个单词,将其字符串编码为 UTF-8 字节序列。Python 中可以通过 word.encode('utf-8') 得到 bytes 对象。

  4. 将每个字节转换为一个可显示的符号。常见的表示方法有:

  5. 直接将字节值(0-255)作为符号 ID,但为了可读性,通常映射为字符形式,如 '!' 对应 33,而不可打印字节则转为形如 <0x00> 的字符串或直接保留为 bytes 对象。
  6. 在 GPT-2 的实现中,字节被映射到 Unicode 字符(例如 0 -> 'Ā', 1 -> 'ā', ...),但逻辑本质不变。
  7. 为了简单,我们可以使用整数字节值(0-255)或使用 bytes 单字节切片作为符号,但字典键需可哈希。常用做法是将字节转换为整数,但在符号序列中使用整数,便于合并。

  8. 在每个单词字节序列末尾添加一个特殊的结束符,可以是一个特殊字节值(如 256,或者 </w> 的某种表示)。为了不与 0-255 冲突,通常使用超出范围的 ID(如 256 表示结束符)。或者沿用字符串表示,用 '</w>' 字符串符号,但在字节层面,所有符号都是整数,可定义 end_of_word = 256

  9. 统计每个字节(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 ab 替换为新的 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 更新频率是合理且简单的。

更新步骤(每次合并后):

  1. 选出 best_pair = (a, b),新符号 c = a + b(或新 ID)。

  2. 获取受影响的单词索引集合:affected_indices = pair_to_word_indices.get(best_pair, set()),如果为空则忽略(理论上不会)。

  3. 对于每个索引 idxaffected_indices 中:

  4. 取出 symbols, freq = word_symbols[idx]
  5. 减去旧 bigram 频率:扫描旧 symbols 的所有相邻对,将其频率从 bigram_freq 中减去 freq。同时从 pair_to_word_indices[pair] 中移除该索引(如果该单词不再包含该 pair,则移除;由于我们后面会重新添加新 pair,可以一次性清除:删除所有旧 pair 的关联,之后重新建立)。
  6. 应用合并:在 symbols 上执行 (a,b) -> c 的合并(同之前的 merge_vocab 逻辑),得到 new_symbols
  7. 添加新 bigram 频率:扫描 new_symbols 的所有相邻对,将其频率加上 freq。同时将该索引 idx 加入到对应 pair_to_word_indices[pair] 集合中。
  8. 更新 word_symbols[idx] 的序列为 new_symbols

  9. best_pair 的频率清零并从 pair_to_word_indices 中删除,因为所有出现都被合并了(除非合并后由于某种原因重新形成,但在这个局部更新中,旧 pair 已经被消除,新序列中不会包含 (a,b) 除非 c 再参与其他合并,但 (a,b) 作为原始 pair 不会再出现,因为 ab 相邻的情况被替换了。注意:如果序列为 ... a a b ...,合并时扫描从左到右只会替换第一个匹配项?需要确保扫描策略与全量重建一致:一次扫描替换所有非重叠的 (a,b)。那么合并后可能仍有剩余的 ab 相邻吗?例如 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 是安全的。

  10. (可选)如果新产生的 bigram 在 bigram_freq 中不存在,则创建;如果某个 bigram 频率降为 0,则从 bigram_freqpair_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 中记录的一致,不一致则丢弃并重新弹出。这种方法实现简单且常用。

  • 主动更新:在堆中直接修改元素,但需要额外索引,复杂且不必要。

懒删除策略:

  1. 维护 bigram_freq 字典作为实时频率的权威源。

  2. 初始化:将所有 (pair, freq)(-freq, pair) 形式压入列表,然后 heapify

  3. 每次需要选择最佳 pair 时:

  4. 从堆中弹出堆顶元素,检查其频率(取负)是否等于 bigram_freq[pair] 的当前值。
  5. 如果相等,说明是有效的最高频 pair;否则丢弃,继续弹出,直到找到有效的。
  6. 需要注意,堆可能为空(所有 pair 频率为 0 或无 bigram)。

  7. 选定 best_pair 后,执行合并,并使用上一问的局部更新方式更新 bigram_freq

  8. 更新过程中,某些 pair 的频率增加或减少,甚至新增 pair。对于这些变化的 pair,我们将 新的 (-new_freq, pair) 推入堆中,不删除旧条目。旧条目会在未来被弹出时因频率不匹配而丢弃。

  9. 为了避免堆无限增长,可以定期清理,但通常合并次数有限,堆大小可控。

堆中元素格式:由于可能出现频率相同的情况,还需要实现稳定的平局打破。如果仅存储 (-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 分词器,生成 mergesvocab(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 分为 1990s)。这可以确保标点成为独立的“单词”,并加上结束符,使得 BPE 能够学习到标点后通常跟随词尾等特点。

  • 对于数字,可以将连续数字作为一个整体单词,也可以按位拆分。但通常保留连续数字作为整体,让 BPE 内部按需合并(如 100 可能合并为 1 0 0 或者整体)。如果语料中有很多数字,BPE 可能学习到 100 的整体表示,但更好的做法是让数字保持为字符级或字节级,由 BPE 自动处理。

  • 预处理伪代码:在训练前加入一个分词器(tokenizer),将原始文本拆分成单词序列(每个单词将独立加上 </w>)。这个分词器的行为是:按空格和标点拆分,保留标点作为单独 token。

具体预处理步骤:

  1. 利用正则表达式将文本分割为 tokens(单词和标点独立)。例如,模式 r'\w+|[^\w\s]' 可以匹配所有单词字符序列或者非空白非单词字符(标点)。也可以处理缩写中的标点如 's 等。

  2. 对于每个 token,统计频率,将它们作为“单词”输入 BPE 初始化。即 token 序列中的每一项都看作一个单词,末尾加结束符。

  3. 对于数字,可以选择保留连续数字,或者进一步拆分为单个数字。但 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,编码时如果用户要求,直接加入输出,而不经过字符拆分与合并。

实现步骤:

  1. 定义特殊 token 列表,如 special_tokens = ["[PAD]", "[UNK]", "[CLS]", "[SEP]", "[MASK]"]

  2. 在词频统计阶段,为每个特殊 token 分配一个极小的频率(或 0),使得它们出现在词汇表中。或者可以在初始化词汇表时直接插入它们。

  3. 初始化符号序列时,对于特殊 token 单词,不进行字符拆分,而是直接将其整个字符串作为一个符号,并加上 </w>(或直接用自身作为词尾,免去结束符)。例如,"[CLS]" -> ['[CLS]', '</w>']

  4. 在 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.jsonmerges.txt(或一并打包在 tokenizer.json 中)。vocab.json 记录了 token 到 ID 的映射,merges.txt 记录了合并规则,每行是一对 token(用空格分隔),按顺序排列(最早合并的在第一行)。我们需要解析这些文件并构建编码/解码逻辑。

解析步骤:

  1. 读取 merges.txt:文件头可能包含版本注释,真正的合并规则从某一行开始。标准格式是 #version: 0.2 然后换行之后每行 token_a token_b。我们将这些行解析为合并列表,每项为 (token_a, token_b)。合并顺序即为文件中的出现顺序。

  2. 读取 vocab.json:得到 {token: id} 字典。同时,它也包含了所有基础字符和合并产生的子词。注意这个 vocab 的 ID 可能不连续或有特殊 token。我们需要反向映射 id_to_token 用于解码。

  3. 构建 token 到 ID 映射:直接用 vocab.json 的字典。

  4. 编码:与之前类似,输入单词,拆分为字符序列,加上 </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 字段里的 vocabmerges(可能是数组)。可以类似解析:

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=15num_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 分词器。