从并行 BPE 训练到手写 Tokenizer:一次围绕类型、特殊 token 和流式编码的实现记录
这篇文章记录的是完成 CS336 Assignment 1 中 BPE 并行训练和
Tokenizer类的过程。它不是直接贴最终答案,而是按真实实现顺序整理:先把train_bpe的预分词阶段尝试并行化,再手写Tokenizer.encode、decode和encode_iterable。整个过程里最容易卡住的不是 BPE 核心算法本身,而是类型表示、special token 的处理顺序、以及流式输入的边界问题。
前一篇文章已经记录了 train_bpe 主流程:从 Karpathy 的最小 BPE demo 出发,逐步适配 CS336 的 bytes-level vocab、GPT-2 预分词、special tokens、Counter 去重、pair_to_token 增量更新和 heap 优化。
这篇文章接着往后写两个问题:
- 第一,
train_bpe的预分词阶段怎么并行化; - 第二,训练得到的
vocab和merges怎么放进一个真正能encode/decode的Tokenizer类里。
1. BPE 并行化训练
1.1 为什么会想到并行化:merge 已经快了,慢点转移到了预分词
在 train_bpe 里,最开始的瓶颈是 merge 循环。朴素版本每一轮都扫描所有 token,后来通过 Counter 把重复 pre-token 合并,再用 pair_to_token 只更新受影响的 token,最后又用 heap 减少 max(stats) 的全量扫描。
做到这一步之后,merge 阶段已经明显变快,新的瓶颈开始转移到前面:
text
读取文本
按 special token 切分
对普通文本跑 GPT-2 regex
把 pre-token encode 成 bytes tuple
统计 Counter
这些操作里,尤其是 GPT-2 regex 和 str.encode("utf-8"),在大语料上会重复执行很多次。于是自然想到:既然每个文本 chunk 的预处理相互独立,能不能把这部分拆给多个进程做?
1.2 第一个并行版本:AsyncResult 不是结果本身
一开始的并行思路是:
text
把文件切成多个 chunk
每个 chunk 交给一个子进程做 init_text
主进程把所有结果收集起来
再统一 Counter
这里用到了 multiprocessing.Pool.apply_async。当时第一个容易误解的点是:apply_async 返回的不是子进程算出来的文本结果,而是一个 AsyncResult 句柄。
也就是说:
text
results 里装的不是 list[str]
results 里装的是 AsyncResult
所以不能直接:
python
for texts in results:
for text in texts:
...
因为 texts 此时还不是可迭代的真实结果。必须先调用:
python
texts.get()
get() 的含义是:等待子进程完成,并把子进程返回的对象取回到主进程里。
这里还讨论过一个误区:AsyncResult 不是"存了一个子进程对象地址"。多进程之间内存空间是隔离的,子进程返回结果时,会经过序列化传回主进程。Pool 关闭也不是说"地址消失所以不能 get",而是如果任务已经正常返回,主进程可以通过 AsyncResult.get() 拿到序列化后的结果;如果没有 get(),那你手里就一直只是任务句柄,不是真实对象。
1.3 子进程应该返回什么:从 list[str] 改成 Counter
第一版并行方案是让子进程返回 init_text 后的 text_list,主进程再把所有文本聚合起来,然后统一:
python
Counter(tuple(text.encode("utf-8")) for text in new_text_list)
这个版本能理解,但不是最好的工程拆分。因为如果子进程已经拿到了自己的 chunk,它完全可以在子进程内部完成:
text
init_text
encode 成 bytes tuple
Counter 计数
然后主进程只需要合并多个 Counter。
这样做有两个好处:
- 子进程返回的数据更紧凑,不需要把大量重复 pre-token 原样传回主进程;
- 和原来
train_bpe的Counter优化保持一致,后续 merge 阶段仍然接收token -> count的结构。
这里也顺便把 Counter 的理解理清了:Counter 本质上是一个字典,key 是 token,value 是出现次数。同一个 key 会自动合并计数,所以它看起来有点像"自动去重",但更准确地说,它是"按 key 计数"。
例如:
text
("t", "h", "e") -> 100
("a", "n", "d") -> 80
多个子进程返回 Counter 之后,主进程可以从空 Counter 开始累加。之前还讨论过 Counter 的 + 会丢掉 0 或负数计数,这在训练预分词统计阶段通常不是问题,因为这里的 count 都是正数;但如果以后写差分统计或 subtract,就要注意这个语义。
1.4 Windows 多进程和 main 入口保护
并行版本还有一个 Python 工程问题:多进程入口保护。
在 Windows 上,multiprocessing 默认使用 spawn 方式启动子进程。子进程会重新导入当前脚本。如果脚本顶层直接启动 Pool,就可能出现子进程导入脚本时又启动新的子进程,递归创建进程。
所以训练脚本应该写成:
text
定义 task
定义 run_train_bpe_parallel
定义 main
if __name__ == "__main__":
main()
当时还讨论了怎么给脚本传参。这里用的是 argparse。可以把 argparse 理解成 Python 标准库里专门处理命令行参数的工具箱,而 argparse.ArgumentParser(...) 是从这个工具箱里创建出来的"参数解析器对象"。
它负责三件事:
- 声明脚本需要哪些参数;
- 从命令行读取这些参数;
- 把字符串参数转换成需要的类型,比如
int。
这部分不是 BPE 算法本身,但它是把并行训练函数做成可执行脚本时必须处理的工程细节。
2. 手写 Tokenizer 类
2.1 写 Tokenizer 的真正难点:输入是 str,内部却是 byte-level BPE
进入 Tokenizer 类之后,问题变得和 train_bpe 不一样。
训练时我们已经得到了:
python
vocab: dict[int, bytes]
merges: list[tuple[bytes, bytes]]
也就是说,词汇表是:
text
token id -> bytes
合并表是:
text
(bytes, bytes) -> merged bytes
但 encode 的输入是:
text
str
这就带来第一个大坑:类型必须对齐。
BPE merge 表描述的是 byte-level 的合并规则,比如:
text
(b"a", b"b") 合并成 b"ab"
而不是:
text
("a", "b") 合并
也不是:
text
(97, 98) 合并
所以 encode 不能一直在 str 层面操作,也不能随便把普通文本转成 int 后再和 bytes merge 表比较。它需要先经过 GPT-2 预分词,再把普通 pre-token 转成 bytes 单元,后续 merge 才能和 merges 里的 bytes pair 对上。
因此初始化时要保留两个方向的 vocab:
text
self.vocab: id -> bytes,用于 decode
self.rev_vocab: bytes -> id,用于 encode 最后反查 id
这个反查表非常关键。因为 encode 的最后一步不是生成 bytes,而是生成 token id。
2.2 类型转换问题:bytes 一遍历就会变成 int
实现 encode 时,最容易踩的 Python 细节是:
python
for x in b"abc":
...
这里的 x 不是 b"a"、b"b"、b"c",而是:
text
97, 98, 99
也就是说,遍历 bytes 得到的是整数。
这正是前面类型不匹配的根因。普通文本如果写成:
python
for byte in text_bytes:
token_list.append(byte)
那 token_list 里放进去的是 int;但 merges 里的 pair 是 bytes。后面判断:
text
token[i] == old_idx1
就会变成:
text
97 == b"a"
当然匹配不上。
解决方式是把单个 byte value 再包回 bytes:
python
bytes([byte])
这样普通文本的内部表示才会变成:
text
[b"a", b"b", b"c"]
这一步解决的是 str -> bytes -> list[bytes] 的类型链路。
2.3 special token 必须一开始就识别,不能最后再补救
第二个大问题是 special token。
special token 的本质是"整体保留"。比如:
text
<|endoftext|>
它不能先被拆成普通字节,再在后面尝试识别。因为一旦拆成:
text
b"<", b"|", b"e", ...
后续 BPE merge 可能会把其中一部分和普通字符合并,原来的整体边界信息就丢了。
所以 special token 的识别必须发生在最前面:
text
原始字符串
先识别 special token
special token 直接转成 token id
普通文本再进入 GPT-2 regex 和 BPE merge
这也是为什么 encode 内部最后形成的是一个混合结构:
text
special token: int
普通 pre-token: list[bytes]
例如:
text
[ [b"H", b"e", b"l", b"l", b"o"], 50256, [b" ", b"w", b"o", b"r", b"l", b"d"] ]
这里的 50256 表示 special token 已经被直接转成 id,不再参与普通 BPE merge。
2.4 findall 和 split 的区别:为什么 special token 要用保留分隔符的切法
一开始我们也讨论过直接把 special token pattern 和 GPT-2 pattern 拼成一个大正则,然后 findall。
这种方式有时能跑,但语义不够清楚。后来更稳的方向是:
text
先按 special token 切原文,并保留 special token 本身
再对普通片段跑 GPT-2 regex
这里关键是 re.split 要带捕获组:
text
(special_token_1|special_token_2|...)
带捕获组时,re.split 会把分隔符本身也放回结果里。这样才能区分:
text
普通文本片段
special token 片段
普通文本片段
如果只用 findall 匹配 special token pattern,它只会返回匹配到的 special token,而不会返回中间那些普通文本。这个点当时很容易误解,因为 findall 看起来像"正则化切分",但它本质上是"找出所有匹配项",不是"按规则切开并保留剩余文本"。
另外,special tokens 需要按长度从长到短排序。原因是可能有重叠 special token:
text
<|endoftext|>
<|endoftext|><|endoftext|>
如果短的排在前面,长 special token 可能先被切成两个短 special token,测试就会失败。长的优先,才能保证"最长 special token 优先匹配"。
2.5 encode 里的 pair_to_token:从训练阶段迁移过来,但结构要变
写完基础版本后,又把 train_bpe 里的 pair_to_token 思路迁移到了 encode。
在 train_bpe 里,pair_to_token 是:
text
pair -> set[token_tuple]
因为训练阶段的数据结构是:
text
unique_token_list: Counter[token_tuple -> count]
token 是不可变 tuple,可以作为字典 key,也能安全放进 set。
但在 encode 里,情况不一样。这里的 token_list 是一个有顺序的列表,每个普通 pre-token 是 list[bytes],special token 是 int。普通 token 后续会被不断 merge 并替换,所以直接把 list 放进 set 不行,list 本身也不可哈希。
因此 encode 里的反向索引更适合写成:
text
pair -> set[token_index]
也就是记录某个 pair 出现在哪些 token_list 下标里。
这样每轮 merge 时,只处理受影响的 token:
text
找到包含 merged_pair 的 token 下标
重写这些 token
删除旧 pair 的索引
加入新 pair 的索引
这个版本相比朴素 encode 的优势是:不用每条 merge rule 都扫描所有 pre-token。尤其 GPT-2 的 merge 表很长,如果每一轮都全量扫,速度会明显变慢。
这里也有几个实现细节:
- 遍历
pair_to_token[merged_pair]时最好先转成list(...),因为循环过程中会修改pair_to_token; - 删除旧 pair 时用
discard,因为元素不存在也不会报错; - 如果
pair已经被删掉,访问defaultdict(set)可能会重新创建空 set,所以删除前可以先判断if pair not in pair_to_token; - 普通 token 合并后生成的中间 token 仍然用 bytes 表示,比如
b"".join(merged_pair),最后再通过rev_vocab反查 id。
这就是 pair_to_token 在训练和 encode 里的核心差异:
text
train_bpe:
pair -> token_tuple
目标是更新 Counter、stats、pair_to_token
encode:
pair -> token_list index
目标是更新当前输入文本的 token_list
2.6 最后一步 flatten:不能把整个 list[bytes] 直接查 vocab
encode 的最后输出必须是:
text
list[int]
但内部的 token_list 是混合结构:
text
int
list[bytes]
所以最后 flatten 时要分情况:
text
如果 token 是 int,说明它是 special token id,直接 append
如果 token 是 list[bytes],就遍历里面每个 bytes,再用 rev_vocab 查 id
当时一个容易错的地方是想直接对整个 token 做反查:
text
rev_vocab[token]
但此时 token 是一个列表,比如:
text
[b"He", b"llo"]
它不是单个 bytes,也不能作为 rev_vocab 的 key。真正应该反查的是列表里的每一个 bytes token:
text
b"He" -> id
b"llo" -> id
这一步把前面的内部 byte-level 表示最终转换回作业要求的 token id。
2.7 decode 反而很简单:先拼 bytes,再统一 UTF-8 解码
相比 encode,decode 简单很多。
因为 vocab 本来就是:
text
id -> bytes
所以 decode 只需要:
text
根据 id 找到 bytes
把所有 bytes 拼起来
最后整体 decode 成 str
这里重要的是"整体 decode",而不是每个 token 单独 decode。原因是 Unicode 字符可能由多个 bytes 组成,如果每个 token 单独 decode,可能会把一个字符拆坏。
所以正确思路是:
text
ids -> bytes list -> b"".join(...) -> decode("utf-8", errors="replace")
errors="replace" 的作用是:遇到非法 UTF-8 字节序列时,不直接抛异常,而是用替换字符处理。这和 tokenizer 测试里的鲁棒性要求更匹配。
3. 流式编码 encode_iterable
3.1 真正麻烦的是流式边界
最后是 encode_iterable。
这个函数的接口是:
python
encode_iterable(self, iterable: Iterable[str]) -> Iterator[int]
它的输入不是一个完整字符串,而是一段一段来的字符串;输出也不是一次性返回列表,而是通过 yield 一个一个产出 token id。
最开始最直接的想法是:对 iterable 里的每个 str 直接调用 encode,然后把结果逐个 yield 出去。这种写法最简单,也最省内存,但它隐含了一个前提:每个 chunk 的结尾都是安全边界。实际并不一定。比如 special token 是 <|endoftext|>,输入被拆成 "<|endo" 和 "ftext|>",第一段就会被当成普通文本编码;等第二段来了,已经没法再把两段合成一个完整 special token。普通 pre-token 也有类似问题:如果 chunk 恰好断在某个 merge pair 的两个 token 中间,那么这两个 token 就无法在 BPE 阶段完成合并,最终编码结果就可能和对完整文本一次性 encode 的结果不同。
然后考虑过基于长度的"保留尾巴"方案。思路是维护一个 buffer,每次只处理前面确定安全的部分,末尾留下一段暂时不编码。这里同时考虑两种 unsafe 边界:一种是 special token 可能被截断,所以根据 max_special_token_len - 1 保留尾部;另一种是 GPT-2 pre-token 可能在 chunk 末尾还没结束,所以找到最后一个可能未完成的 pre-token 起点。最终取两个 unsafe 边界里更靠前的那个,也就是 min(special_unsafe_start, gpt2_unsafe_start),尽量避免既切断 special token,又切断普通 pre-token。
但这个方案仍然不是严格正确。因为它本质上还是靠长度和位置做保守截断,不能真正判断末尾是否已经稳定。special token 之间可能有前缀关系,比如同时存在 "<|end|>" 和 "<|end|> ",当 buffer 末尾是 "<|end|>" 时,它虽然已经是完整 special token,但如果下一个 chunk 开头是空格,就应该匹配更长的 special token。GPT-2 pre-token 也有类似问题,连续字母、数字、符号、空白都可能继续延长,所以固定长度或简单保留最后一个 pre-token 都无法从根本上证明边界安全。
后来又考虑过更成熟的通用流式方案:先把 special token 的处理放在 GPT-2 regex 之前,用前缀表或 trie 检查 buffer 末尾是否可能是某个 special token 的前缀;普通文本再用 GPT-2 regex 做 pre-token 切分,并保留最后一个还不确定的 pre-token。整体思路是"只 encode 稳定前缀,保留不稳定后缀,等下一段输入来了再继续判断"。这个方案理论上更接近严格流式 tokenizer,但实现复杂度明显上升,因为 special token 的不稳定后缀和 GPT-2 pre-token 的不稳定后缀可能重叠,边界判断很容易写错。
最后结合测试场景做了工程取舍:测试里 encode_iterable 主要是接收文件对象,而文件对象迭代通常是按行读入。于是最终选择面向测试的实现方式:对每个输入片段直接调用现有 encode,然后逐个 yield 里面的 id。这个方案不是任意 chunk 切分下都严格正确的通用流式 tokenizer,但它简单、低内存,复用了已经写好的 encode,也符合当前测试输入方式。
真实工程里也不一定非要让 tokenizer 适配任意切分方式的流输入。另一种更优解可能是反过来限制输入协议,比如要求上游按行、按文档块、按已知安全边界输入,或者保证 special token 不会被拆开。这样可以把复杂的边界恢复逻辑从 tokenizer 内部移走,用清晰的输入约束换取更简单、更稳定、更容易验证的实现。
4. 总结
4.1 这次手写 tokenizer 真正解决的不是一个问题,而是一串适配问题
回头看这次过程,Tokenizer 难的地方并不是 BPE 的概念本身。BPE 的核心仍然是:
text
统计 pair
按 merges 顺序合并
输出 token id
真正复杂的是把这个核心逻辑放进 CS336 的接口和测试语义里。
第一,train_bpe 并行化时,要先想清楚进程之间传什么。直接传大量 pre-token 列表能跑,但不够紧凑;让子进程直接返回 Counter,主进程再合并,是更自然的拆分。这个过程也顺便弄清楚了 AsyncResult.get()、Pool 生命周期、Counter 聚合和 Windows 多进程入口保护。
第二,encode 的核心问题是类型适配。输入是 str,但 vocab 和 merges 都是 byte-level:vocab 是 int -> bytes,merges 是 tuple[bytes, bytes]。所以普通文本必须转成 list[bytes] 后才能参与 merge,最后再通过 rev_vocab 转成 id。中间只要混进 int、str、listbytes 的错误使用,merge 就可能完全匹配不上。
第三,special token 必须在最开始识别。它不是普通文本的一部分,而是一个整体 token。如果先拆成字节,再试图在最后识别 special token,边界信息已经丢了,后续 merge 也可能改变它的内部结构。
第四,pair_to_token 的思路可以从 train_bpe 迁移到 encode,但不能照搬。训练阶段存的是不可变 token tuple,因为要维护 Counter、stats 和 pair_to_token;encode 阶段更适合存 token 下标,因为当前输入的 token_list 会被原地更新。
第五,decode 很直接,但也有一个原则:先把所有 id 对应的 bytes 拼起来,再整体 UTF-8 decode。不要逐 token 解码,否则可能拆坏多字节 Unicode 字符。
第六,encode_iterable 暴露的是流式边界问题。严格通用方案需要维护 buffer、识别 stable prefix、保留 unsafe suffix;但在当前测试输入按文件行迭代的前提下,直接逐段 encode 再 yield from 是一个合理的工程取舍。
所以这次实现最终可以概括成一句话:
text
train_bpe 解决的是如何训练出符合 CS336 要求的 byte-level vocab 和 merges;
Tokenizer 解决的是如何把 str 输入、安全的 special token 处理、byte-level merge 规则和最终 token id 输出接成一条完整链路。
python
class Tokenizer:
def __init__(self,vocab,merge_dic,special_tokens=None):
self.vocab = vocab
self.rev_vocab = {v:k for k,v in vocab.items()}
# 翻转 vocab ,方便后续 encode 从 bytes 找到 token_id
self.merge_dic = merge_dic
self.special_tokens = special_tokens or []
if self.special_tokens:
self.special_tokens.sort(key=len,reverse=True)
self.set_special_tokens = set(self.special_tokens)
new_st = []
for st in self.special_tokens:
st = re.escape(st)
new_st.append(st)
self.split_parten = ("("+"|".join(new_st)+")") if new_st else []
self.gpt2pat = re.compile(r"""'s|'t|'re|'ve|'m|'ll|'d| ?\p{L}+| ?\p{N}+| ?[^\s\p{L}\p{N}]+|\s+(?!\S)|\s+""")
@classmethod
def from_files(cls,vocab_filepath,merges_filepath,special_tokens=None):
with open(vocab_filepath,"r",encoding = "utf-8") as f:
vocab0 = json.load(f)
vocab = {token:bytes(bytes_list) for token,bytes_list in enumerate(vocab0)}
with open(merges_filepath,"r",encoding = "utf-8") as f:
merge0 = json.load(f)
merge = [(bytes(left),bytes(right)) for left,right in merge0]
return cls(vocab,merge,special_tokens)
"""
内存里
• vocab: dict[int, bytes]
• merges: list[tuple[bytes, bytes]]
文件里
• vocab.json 用 list[list[int]],下标就是 token id
• merges.json 用 list[[list[int], list[int]]]
• special_tokens 单独存/单独传,类型用 list[str]
例如:
vocab:
[
[0],
[1],
[97],
[97, 98]
]
merges:
[
[[97], [98]],
[[97, 98], [99]]
]
"""
def encode(self,text):
if self.split_parten:
text_list = re.split(self.split_parten,text)
new_text_list = []
for text in text_list:
if text in self.set_special_tokens:
new_text_list.append(self.rev_vocab[text.encode("utf-8")])
continue
sub_text_list = [text0.encode("utf-8") for text0 in re.findall(self.gpt2pat,text)]
new_text_list.extend(sub_text_list)
else:
new_text_list = [text0.encode("utf-8") for text0 in re.findall(self.gpt2pat,text)]
text_list = new_text_list
# 正则化
token_list = []
for text in text_list:
if isinstance(text,int):
token_list.append(text)
continue
new_token = []
for text0 in text:
new_token.append(bytes([text0]))
token_list.append(new_token)
# 把特殊字符串直接转成 token_id
pair_to_token = defaultdict(set)
for idx,token in enumerate(token_list):
if isinstance(token,int):
continue
for pair in zip(token,token[1:]):
pair_to_token[pair].add(idx)
for merged_pair in self.merge_dic:
old_idx1 = merged_pair[0]
old_idx2 = merged_pair[1]
new_bytes = b"".join(merged_pair)
if merged_pair not in pair_to_token:
continue
for affected_token_idx in list(pair_to_token[merged_pair]):
token = token_list[affected_token_idx]
new_token = []
i = 0
while(i<len(token)):
if(i<len(token)-1 and token[i] == old_idx1 and token[i+1] == old_idx2):
new_token.append(new_bytes)
i += 2
else:
new_token.append(token[i])
i += 1
token_list[affected_token_idx] = new_token
for pair in zip(token,token[1:]):
if pair not in pair_to_token:
continue
pair_to_token[pair].discard(affected_token_idx)
if not pair_to_token[pair]:
del pair_to_token[pair]
for pair in zip(new_token,new_token[1:]):
pair_to_token[pair].add(affected_token_idx)
fin_token = []
for token in token_list:
if isinstance(token, int):
fin_token.append(token)
continue
for byte in token:
fin_token.append(self.rev_vocab[byte])
return fin_token
def decode(self,ids):
text_bytes = b"".join(self.vocab[idx] for idx in ids)
return text_bytes.decode("utf-8",errors="replace")
def encode_iterable(self,iterable):
for text in iterable:
token = self.encode(text)
yield from token