Problem (train_bpe)
dnnlpy 状态:已在 dnnlpy.tokenizers.BPETrainer 中实现。
首先,实现一个最简单的 BPE。这个版本只关注算法流程,不考虑任何优化。在每一次 merge 操作中,直接使用 max() 函数遍历所有候选 pair,选择当前频率最高的 pair。同时,这个实现没有维护 pair 的倒查表,也没有使用多进程、多线程或其他缓存机制。因此,在合并一个 pair 后,需要重新遍历整个词表,找到所有包含该 pair 的 word,并更新对应的统计信息。这种实现虽然简单直观,但大量重复扫描会带来较高的计算开销。
假设进行 \(M\) 次 merge,每次需要扫描 \(P\) 个候选 pair,并遍历 \(V\) 个 word 更新统计,每个 word 的平均长度为 \(L\) ,那么整体计算的复杂度约为:
\[
O(M \cdot (P + VL))
\]
这在大规模语料和较大词表的情况下,训练时间是不可接受的。因此,需要引入一系列优化策略来提升训练效率。
为了确定这个简单实现的性能瓶颈,首先使用 Python 自带的 cProfile 工具进行 profiling。相比直接分析代码哪里慢,profiling 可以统计每个函数的调用次数和累计运行时间,从而定位真正影响性能的部分。
例如,可以通过:
python -m cProfile -s cumulative train_bpe.py
对 BPE 训练过程进行分析。其中:
ncalls 表示函数被调用的次数;
tottime 表示函数自身的运行时间,不包括调用子函数的时间;
cumtime 表示函数及其子函数的累计运行时间。
运行 profile 后,可以发现大量时间消耗在以下几个步骤:
寻找最高频 pair。每一次 merge 都需要调用 max() 遍历所有候选 pair,当 pair 数量较大时,这个过程会被重复执行很多次,导致性能瓶颈。
更新 pair 统计。由于没有维护倒查表,每次 merge 后都需要重新扫描整个词表,寻找包含目标 pair 的 word,并重新计算相关频率。这个步骤通常会占据大量运行时间,因为它涉及大量重复的数据遍历。
因此,总体优化思路是,首先优化 pair 查找过程,不再每次使用 max() 遍历所有 pair,而是维护一个能够快速获取最高频 pair 的数据结构。随后,再进一步优化 merge 后的更新过程,通过保存 pair 到 word 的倒查关系,只更新受到影响的部分。
第一轮优化:使用 heap 加速最高频 pair 查找
通过 profiling 可以发现,简单实现中一个明显的瓶颈是:每一次 merge 都需要调用 max() 遍历所有候选 pair,寻找当前频率最高的 pair。随着 pair 数量增加,这一步会被重复执行大量次数。实际上,在整个 BPE 训练过程中,很多时间都花在了反复寻找最大值上,而不是 merge 本身。
因此,第一步优化是引入 heap(优先队列)。相比每次从全部 pair 中搜索最大值,heap 可以提前维护当前频率最高的 pair,使得获取最高频 pair 的操作更加高效。具体来说,将 pair 按照 frequency 分桶,并维护一个 max-heap。每次需要选择下一个 merge pair 时,直接从 heap 顶部获取当前频率最高的 pair,而不需要重新扫描所有候选 pair。
但是,在 merge 过程中,pair 的 frequency 会不断变化。例如,合并一个 pair 后,相关 word 中的其他 pair 频率可能增加或减少。如果每次更新时都立即在 heap 中寻找并删除旧的记录,删除操作本身仍然需要遍历 heap,反而引入新的开销(我们知道,在 heap 中删除任意元素的最坏时间复杂度是 \(O(N)\) ,因为需要先找到该元素的位置)。
因此,这里采用延迟删除(lazy deletion) 。也就是说,当一个 pair 的频率变化时,不立即删除 heap 中的旧记录,而是直接加入新的记录。这样 heap 中可能暂时存着一些已经过期的数据。当之后从 heap 顶部取出一个 pair 时,再检查它的 frequency 是否仍然是最新的。如果不是,就说明这个记录已经过期,直接丢弃,继续查看下一个 pair。这样,我们可以把删除操作的开销分摊到后续的 pop 操作中,时间复杂度仍然保持在 \(O(\log N)\) 。
Python 3.14 已经提供了原生的 max-heap API,例如 heapify_max、heappush_max 和 heappop_max。在旧版本 Python 中,可以通过保存负 frequency 的方式,使用普通 min-heap 模拟 max-heap。
经过这一轮优化后,pair 查找从每次 merge 都扫描全部 pair,变成了通过 heap 快速获取候选 pair。接下来继续使用 profiling,可以观察新的瓶颈,并进一步优化 pair 更新过程。
第二轮优化:只更新受 merge 影响的 word
第一轮优化解决了如何快速找到最高频 pair 的问题。但是 profiling 后可以发现,训练过程中的另一个主要瓶颈仍然存在:每次 merge 后,都需要重新扫描大量 word 来更新 pair frequency。实际上,一次 merge 只会影响包含当前 pair 的 word。对于不包含这个 pair 的 word,它们的 token 序列完全没有变化,对应的 pair frequency 也不会发生改变。因此,没有必要在每轮 merge 后重新遍历整个语料。
为了解决这个问题,在初始化阶段先完整统计一次 pair 信息,并额外维护 pair 到 word 的反向索引:
pair_counts[pair] -> frequency:保存每个 pair 当前的总频率;
pair_indices[pair] -> set[word_index]:记录哪些 word 包含该 pair。
当选择出当前需要 merge 的 best_pair 后,可以通过 pair_indices[best_pair] 直接找到所有受影响的 word,只对这些 word 执行 merge 和统计更新。
对于每个受影响的 word,只需要比较 merge 前后的局部变化,并更新发生改变的 pair frequency。假设某个 pair 在一个 word 中出现次数变化为:
\[
\Delta n(p) = \eta_{\mathrm{new}}(p) - \eta_{\mathrm{old}}(p)
\]
那么它对整体 frequency 的贡献变化为:
\[
\Delta c(p) = \Delta n(p) \times \mathrm{freq}(\mathrm{word})
\]
只有发生变化的 pair 才需要更新对应的统计信息。
通过这种方式,merge loop 从原来的“每轮扫描整个语料,重新统计所有 pair”变成了“只访问包含当前 merge pair 的 word,并增量更新受到影响的 pair”。虽然需要额外的内存来保存 pair 到 word 的倒查索引,但显著减少了重复扫描的开销。
第三轮优化:优化训练阶段的预分词过程
前两轮优化主要减少了 BPE merge loop 中存在的重复搜索和重复统计,但是 profiling 后仍然可以发现,训练开始之前的 pre-tokenization 阶段存在较大的额外开销。
最初的实现直接调用普通的 pre_tokenize() 流程。这个流程主要是为实际编码服务的,因此会生成完整的 (token, offset) 信息,并在每次遇到 token 时执行对应的 Unicode 转换。但是,对于 BPE 训练来说,我们只需要知道每个 token 出现的次数,并不需要保存 offset 信息。同时,在大规模语料中,相同的 token 往往会重复出现,如果每次都重新执行 Unicode 转换,会产生大量重复计算。
因此,我们为训练阶段增加一个专门的 fast path。这个路径不再生成 (token, offset) 对,而是直接统计原始 piece 的出现频率。随后,只需要对不同的 piece 执行一次 Unicode 转换,再将结果用于后续 BPE 训练。这里的 piece 指的是经过预分词(pre-tokenization)后得到的原始文本片段,还没有经过 BPE merge 的 token。
通过这一轮优化,训练阶段避免了两类不必要的工作:
不再计算训练过程中不会使用的 offset 信息;
不再对重复出现的 token 多次执行相同的 Unicode 转换。
相比前两轮针对 merge loop 的优化,这一次优化减少的是数据预处理阶段的额外开销,使整个 BPE 训练流程更加高效。
第四轮优化:批量并行预分词
经过前几轮优化后,merge loop 中的大量重复计算已经减少,但是 profiling 仍然显示,预分词(pre-tokenization)仍然占据了较多运行时间。
考虑到预分词过程中的每个 document 之间相互独立,因此天然适合并行处理。最直接的想法是为每个 document 创建一个独立的任务,交给多个 worker 执行。但是,实际测试发现,如果任务粒度过小,单独提交大量 document 任务会引入额外开销。例如,任务调度、对象序列化以及进程间数据传输都会消耗时间。当单个任务本身很短时,这些开销甚至可能超过并行计算带来的收益。
因此,我们将多个 document 合并成一个 batch,再把 batch 作为并行任务提交。这样可以减少任务数量,提高每个 worker 的工作量,同时降低调度和通信开销。具体来说,预分词阶段按照固定大小划分 batch:
每个 batch 内包含多个 document;
每个 worker 负责处理整个 batch;
不同 batch 之间可以并行统计 token frequency;
最后将各个 worker 返回的统计结果合并。
在实现上,根据 Python 运行环境选择不同的并行方式:
在普通 Python 环境中,由于 GIL 限制,因此使用 ProcessPoolExecutor 创建多个进程执行;
在支持自由线程的 Python 环境中,直接使用 ThreadPoolExecutor,避免进程创建和数据传输的额外开销;
当 num_workers 等于 1 时,直接使用串行路径,避免不必要的并行管理成本。
需要注意的是,并不是所有 BPE 训练过程都可以并行化。预分词只是统计每个 token 的出现频率,各 batch 之间相互独立,因此可以安全并行。而 BPE merge loop 中,每一次 merge 都依赖前一次 merge 更新后的 pair frequency,因此存在严格的数据依赖关系,仍然需要保持串行执行。
通过这一轮优化,主要减少了训练前数据处理阶段的时间开销,使整体训练流程更加接近实际 tokenizer 工程中的实现方式。
BPE 训练调用栈
从作业脚本开始,完整的训练调用栈如下。所有数据采用流式处理,避免一次性加载整个语料到内存中。
train_bpe_tinystories()
├─ datasets.load_dataset(...)
├─ _batch_iterator(dataset)
├─ Tokenizer.train_from_iterator(...)
│ └─ BPE.train_from_iterator(tokenizer, ...)
│ ├─ BPETrainer(...)
│ └─ BPETrainer.train(texts)
│ ├─ _iter_texts(texts)
│ ├─ _count_pre_tokens(texts)
│ │ ├─ Tokenizer._normalize(text)
│ │ └─ Tokenizer._count_pre_tokens(...)
│ │ ├─ it.batched(..., 1024)
│ │ ├─ parallel_map(...)
│ │ └─ ByteLevelPreTokenizer.count_pre_tokens(...)
│ ├─ _init_vocab(word_counts)
│ ├─ _init_pair_counts(word_symbols, word_freqs)
│ ├─ _init_pair_freq_index(pair_counts)
│ ├─ while len(vocab) < vocab_size
│ │ ├─ _select_best_pair(...)
│ │ ├─ pair_indices.pop(best_pair)
│ │ └─ _merge_pair_and_update_pair_counts(...)
│ │ ├─ _merge_pair(...)
│ │ └─ _update_pair_count(...)
│ └─ _save_result(vocab_tokens, merges)
└─ Tokenizer.save(...)
└─ BPE.save(...)
需要注意的是,当前实现与 CS336 参考实现存在一个细节差异。CS336 在 byte 层面执行 BPE,因此 frequency 相同时按照 raw bytes 的字典序进行 tie-break。而当前实现模仿 Hugging Face Tokenizers,先通过 byte-to-Unicode 映射将 bytes 转换为可逆 Unicode 字符,因此内部比较的是 Unicode 字符串的字典序。由于该映射不保证与原始 byte 顺序一致,少数情况下 merge 顺序可能与 CS336 参考实现不同,但不会影响 BPE 的整体流程。
tokenizer = dltk.Tokenizer(
dltk.BPE(),
pre_tokenizer= dltk.ByteLevelPreTokenizer(add_prefix_space= False ),
decoder= dltk.ByteLevelDecoder(),
num_workers= 1 ,
)
tokenizer.train_from_iterator(
['low lower lowest' , 'newer wider' ],
vocab_size= 300 ,
special_tokens= ['[UNK]' , '[PAD]' , '[CLS]' , '[SEP]' ],
initial_alphabet= dltk.ByteLevelPreTokenizer.alphabet(),
)
print ('Vocab size:' , tokenizer.vocab_size)
print ('First five merges:' , tokenizer.model.merges[:5 ])
Vocab size: 277
First five merges: [('w', 'e'), ('l', 'o'), ('Ġ', 'lo'), ('Ġlo', 'we'), ('Ġlowe', 's')]
Problem (train_bpe_tinystories)
当前实验脚本从 Hugging Face roneneldan/TinyStories 读取 2,119,719 条训练样本,设置目标词表大小为 10,000,并记录训练时间和峰值内存使用。
在优化后的 BPE 实现上进行训练,结果如下:
from train_bpe_tinystories import train_bpe_tinystories
tokenizer = train_bpe_tinystories()
longest_token = max (tokenizer.get_vocab(), key= len )
print (f'Tokenizer vocab size: { tokenizer. vocab_size} ' )
print (f'Longest token: { longest_token!r} (length: { len (longest_token)} )' )
Training tokenizer on TinyStories...
Training completed in 42.7828 seconds.
[NOTE] This should be less than 30 minutes for the TinyStories dataset.
Tokenizer vocabulary size: 10000.
LRU cache info: CacheInfo(hits=5, misses=6, maxsize=100000, currsize=6)
Peak memory usage: 2.7093 GB.
[NOTE] This should be less than 30 GB for the TinyStories dataset.
Tokenizer saved to bpe_tinystories.json.
Tokenizer vocab size: 10000
Longest token: 'Ġaccomplishment' (length: 15)
根据作业要求,TinyStories 数据集训练时间应低于 30 分钟,峰值内存应低于 30 GB。当前实现的训练时间约为 16 秒,峰值内存约为 2.6 GB,满足要求。
训练完成后,tokenizer 被保存为 bpe_tinystories.json,包含 10,000 个 token 和对应的 merge rules。按照内部 byte-to-Unicode 字符串长度计算,最长 token 是 Ġaccomplishment,长度为 15。
对实验脚本进行 profiling,分析训练过程中各个阶段的时间开销:
from train_bpe_tinystories import profile_train_bpe_tinystories
profile_train_bpe_tinystories()
Profile 结果显示,当前实现主要的时间开销集中在预分词和 merge 更新阶段。
首先,regex.findall 占据了大量累计时间:
{method 'findall' of '_regex.Pattern' objects}
cumtime: 374.334s
这部分对应 ByteLevel pre-tokenization 过程。虽然单次调用开销较小,但是 TinyStories 包含超过 200 万条文本样本,大量重复调用导致累计时间较高。这说明预分词阶段仍然存在进一步优化空间,例如减少重复 token 的处理,或者通过 batch 并行减少整体等待时间。
第二个明显瓶颈是 pair frequency 的更新:
trainer.py:_merge_pair_and_update_pair_counts
cumtime: 89.856s
trainer.py:_update_pair_count
cumtime: 82.494s
这部分发生在 BPE merge loop 中。虽然当前实现已经使用 heap 快速获取最高频 pair,并通过反向索引只更新受影响的 word,但每次 merge 后仍然需要修改多个 pair 的统计信息。随着 merge 次数增加,这部分增量更新操作成为主要计算开销之一。
另外,可以看到:
utils.py:parallel_map
cumtime: 106.451s
以及 threading 相关函数占据了一定时间。这说明并行预处理本身已经引入额外的任务调度和同步开销,需要进一步关注任务粒度是否合理。如果 batch 太小,线程或进程调度成本可能抵消并行带来的收益。
同时,由于启用了多线程,cProfile 中显示的 cumtime 是所有线程累计执行时间,而不是程序实际经过的 wall-clock time。因此,profile 中某些函数的累计时间可能超过实际训练时间,这是并行执行导致的正常现象。