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产生的vocabmerges,可用于把文本编码为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