BPE Tokenizer 实现与优化
搭建环境(WSL)(并行)
完成 train_bpe 函数
基本逻辑
考虑Unicode Code会导致vocab极大,byte-level tokenization又因编码后的序列过长而减慢模型训练,为了综合两者优劣,即考虑先用初始byte作vocab,每次合并出现频率最大的token pair作为新token加入vocab,直到vocab大小到达指定上限
实际实现过程中,为了划定合并边界并提高合并效率,会先用预分词把训练用文本分为几个pre-token,并拒绝pre-token间的merge
初始实现架构
pre_tokenization(corpus to pretoken_freq) -> count_pairs(pretoken_freq to pair_freq) -> find_max_pair -> merge_pairs
最后用train_bpe函数整合为训练循环
其余实现细节:同频pair取字典序大者
完成 Tokenizer 类(文件加载接口待补充)
这个类用于接收训练好的BPE Tokenizer产生的vocab与merges,可用于把文本编码为token ID,或把token ID解码为文本
结构
def __init__(self, vocab, merges, special_tokens=None)
def from_files(cls, vocab_filepath, merges_filepath, special_tokens=None)
def encode(self, text: str) -> list[int]
def encode_iterable(self, iterable: Iterable[str]) -> Iterator[int]
def decode(self, ids: list[int]) -> str
实现细节:merges顺序需保持
优化 Tokenizer 类(encode)
优化merge逻辑:只需要在一个pre-token中按merges中出现的pair顺序merge即可,不需要顺序枚举merges再扫描每个pre-token
发现
dict查找快于list(键查询与线性扫描的区别)
优化 train_bpe 过程
采用cProfile观察各部分所占时间从而找到bottleneck,采用/usr/bin/time -v观察每次训练的 wall time 与 Max RSS;在多进程场景中,Max RSS 表示最大子进程的峰值常驻内存,并非整个进程池的总内存
优化过程
优化前
每轮重新统计全部pair,并扫描全部pre-token完成合并。
优化 merge_pair
维护一个pair到含有对应pair的pre-token的集合的映射,每次合并pair只需要修改对应pre-token;pre-token的信息维护逻辑修改为pretoken_id与pre-token与其频率的映射,避免因pre-token反复修改而导致字典的修改
优化 find_max_pair 后
采用优先队列(堆)的思路,维护一个heap,在每次合并时把频率发生更改的pair加入heap,寻找max_pair时找到freq最大的且堆中频率与真实频率相同的pair
类中
__lt__用于堆中大小关系构建
cProfile 数据
| 优化阶段 | 函数 | cumtime(s) |
|---|---|---|
| 优化前 | merge_pair |
193.7 |
| 优化前 | count_pairs |
58.7 |
| 优化前 | find_max_pair |
15.2 |
| 优化前 | pre_tokenize |
10.2 |
优化 merge_pair 后 |
find_max_pair |
14.96 |
优化 merge_pair 后 |
pre_tokenize |
10.43 |
优化 merge_pair 后 |
merge_pair |
0.71 |
优化 merge_pair 后 |
count_pairs |
0.14 |
优化 find_max_pair 后 |
pre_tokenize |
10.435 |
优化 find_max_pair 后 |
merge_pair |
0.708 |
优化 find_max_pair 后 |
其他部分 | 0.255 |
/usr/bin/time -v 数据
| 数据集 | vocab |
实现阶段 | wall time | Max RSS(最大子进程,非进程池总内存) | 正确性 |
|---|---|---|---|---|---|
| valid | 10000 | 原始实现 | 99.88s | 94.1MiB | correctness passed |
| valid | 10000 | 优化merge_pair后 |
10.13s | 101.4MiB | correctness passed |
| valid | 10000 | 优化find_max_pair后 |
5.47s | 102.3MiB | correctness passed |
优化 pre_tokenize 后(增加并行逻辑)
导入multiprocessing库中的Pool,实现多进程预分词
| 进程数 | real time 中位数 | user time 中位数 | CPU使用率中位数 | Max RSS 中位数(最大子进程,非进程池总内存) |
|---|---|---|---|---|
| 1 | 5.45s | 4.94s | 94% | 89.5MiB |
| 2 | 3.21s | 5.26s | 167% | 55.4MiB |
| 4 | 2.11s | 5.57s | 272% | 43.9MiB |
| 8 | 1.71s | 7.07s | 432% | 45.0MiB |
| 16 | 1.46s | 9.48s | 660% | 46.5MiB |
从8个进程增加到16个进程,real time中位数仅降低约14.6%,user time中位数增加约34.1%。边际收益有限,因此采用8个进程。
正式训练
使用pickle导入/导出训练结果
TinyStories实验结果
- 训练时间:约 1 分 40.76 秒
/usr/bin/time -v报告的 Max RSS:约 1.48 GiB(最大子进程,非进程池总内存)- 最长 token:
b' accomplishment'(其中之一)
附录:多进程预分词原始测试数据
| 进程数 | 运行次数 | real time | user time | CPU使用率 | Max RSS(最大子进程,非进程池总内存) |
|---|---|---|---|---|---|
| 1 | 1 | 5.45s | 4.93s | 94% | 89.5MiB |
| 1 | 2 | 5.46s | 4.94s | 94% | 89.5MiB |
| 1 | 3 | 5.42s | 5.01s | 95% | 89.4MiB |
| 2 | 1 | 3.21s | 5.26s | 167% | 55.5MiB |
| 2 | 2 | 3.12s | 5.16s | 168% | 55.4MiB |
| 2 | 3 | 3.29s | 5.32s | 165% | 55.4MiB |
| 4 | 1 | 2.32s | 5.75s | 256% | 43.7MiB |
| 4 | 2 | 2.11s | 5.57s | 272% | 43.9MiB |
| 4 | 3 | 2.05s | 5.50s | 276% | 44.1MiB |
| 8 | 1 | 1.71s | 6.93s | 421% | 44.7MiB |
| 8 | 2 | 1.61s | 7.07s | 450% | 45.0MiB |
| 8 | 3 | 1.77s | 7.37s | 432% | 45.1MiB |
| 16 | 1 | 1.65s | 10.16s | 646% | 46.5MiB |
| 16 | 2 | 1.45s | 9.23s | 660% | 46.5MiB |
| 16 | 3 | 1.46s | 9.48s | 675% | 46.6MiB |


