并行 BPE 训练 与 手写 `Tokenizer`
从并行 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 阶段已经明显变快新的瓶颈开始转移到前面读取文本 按 special token 切分 对普通文本跑 GPT-2 regex 把 pre-token encode 成 bytes tuple 统计 Counter这些操作里尤其是 GPT-2 regex 和str.encode(utf-8)在大语料上会重复执行很多次。于是自然想到既然每个文本 chunk 的预处理相互独立能不能把这部分拆给多个进程做1.2 第一个并行版本AsyncResult不是结果本身一开始的并行思路是把文件切成多个 chunk 每个 chunk 交给一个子进程做 init_text 主进程把所有结果收集起来 再统一 Counter这里用到了multiprocessing.Pool.apply_async。当时第一个容易误解的点是apply_async返回的不是子进程算出来的文本结果而是一个AsyncResult句柄。也就是说results 里装的不是 list[str] results 里装的是 AsyncResult所以不能直接fortextsinresults:fortextintexts:...因为texts此时还不是可迭代的真实结果。必须先调用texts.get()get()的含义是等待子进程完成并把子进程返回的对象取回到主进程里。这里还讨论过一个误区AsyncResult不是“存了一个子进程对象地址”。多进程之间内存空间是隔离的子进程返回结果时会经过序列化传回主进程。Pool关闭也不是说“地址消失所以不能 get”而是如果任务已经正常返回主进程可以通过AsyncResult.get()拿到序列化后的结果如果没有get()那你手里就一直只是任务句柄不是真实对象。1.3 子进程应该返回什么从list[str]改成Counter第一版并行方案是让子进程返回init_text后的text_list主进程再把所有文本聚合起来然后统一Counter(tuple(text.encode(utf-8))fortextinnew_text_list)这个版本能理解但不是最好的工程拆分。因为如果子进程已经拿到了自己的 chunk它完全可以在子进程内部完成init_text encode 成 bytes tuple Counter 计数然后主进程只需要合并多个Counter。这样做有两个好处子进程返回的数据更紧凑不需要把大量重复 pre-token 原样传回主进程和原来train_bpe的Counter优化保持一致后续 merge 阶段仍然接收token - count的结构。这里也顺便把Counter的理解理清了Counter本质上是一个字典key 是 tokenvalue 是出现次数。同一个 key 会自动合并计数所以它看起来有点像“自动去重”但更准确地说它是“按 key 计数”。例如(t, h, e) - 100 (a, n, d) - 80多个子进程返回Counter之后主进程可以从空Counter开始累加。之前还讨论过Counter的会丢掉 0 或负数计数这在训练预分词统计阶段通常不是问题因为这里的 count 都是正数但如果以后写差分统计或 subtract就要注意这个语义。1.4 Windows 多进程和main入口保护并行版本还有一个 Python 工程问题多进程入口保护。在 Windows 上multiprocessing默认使用 spawn 方式启动子进程。子进程会重新导入当前脚本。如果脚本顶层直接启动Pool就可能出现子进程导入脚本时又启动新的子进程递归创建进程。所以训练脚本应该写成定义 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不一样。训练时我们已经得到了vocab:dict[int,bytes]merges:list[tuple[bytes,bytes]]也就是说词汇表是token id - bytes合并表是(bytes, bytes) - merged bytes但encode的输入是str这就带来第一个大坑类型必须对齐。BPE merge 表描述的是 byte-level 的合并规则比如(ba, bb) 合并成 bab而不是(a, b) 合并也不是(97, 98) 合并所以encode不能一直在str层面操作也不能随便把普通文本转成 int 后再和 bytes merge 表比较。它需要先经过 GPT-2 预分词再把普通 pre-token 转成bytes单元后续 merge 才能和merges里的 bytes pair 对上。因此初始化时要保留两个方向的 vocabself.vocab: id - bytes用于 decode self.rev_vocab: bytes - id用于 encode 最后反查 id这个反查表非常关键。因为 encode 的最后一步不是生成 bytes而是生成 token id。2.2 类型转换问题bytes一遍历就会变成int实现encode时最容易踩的 Python 细节是forxinbabc:...这里的x不是ba、bb、bc而是97, 98, 99也就是说遍历bytes得到的是整数。这正是前面类型不匹配的根因。普通文本如果写成forbyteintext_bytes:token_list.append(byte)那token_list里放进去的是 int但merges里的 pair 是 bytes。后面判断token[i] old_idx1就会变成97 ba当然匹配不上。解决方式是把单个 byte value 再包回 bytesbytes([byte])这样普通文本的内部表示才会变成[ba, bb, bc]这一步解决的是str - bytes - list[bytes]的类型链路。2.3 special token 必须一开始就识别不能最后再补救第二个大问题是 special token。special token 的本质是“整体保留”。比如|endoftext|它不能先被拆成普通字节再在后面尝试识别。因为一旦拆成b, b|, be, ...后续 BPE merge 可能会把其中一部分和普通字符合并原来的整体边界信息就丢了。所以 special token 的识别必须发生在最前面原始字符串 先识别 special token special token 直接转成 token id 普通文本再进入 GPT-2 regex 和 BPE merge这也是为什么encode内部最后形成的是一个混合结构special token: int 普通 pre-token: list[bytes]例如[ [bH, be, bl, bl, bo], 50256, [b , bw, bo, br, bl, bd] ]这里的50256表示 special token 已经被直接转成 id不再参与普通 BPE merge。2.4findall和split的区别为什么 special token 要用保留分隔符的切法一开始我们也讨论过直接把 special token pattern 和 GPT-2 pattern 拼成一个大正则然后findall。这种方式有时能跑但语义不够清楚。后来更稳的方向是先按 special token 切原文并保留 special token 本身 再对普通片段跑 GPT-2 regex这里关键是re.split要带捕获组(special_token_1|special_token_2|...)带捕获组时re.split会把分隔符本身也放回结果里。这样才能区分普通文本片段 special token 片段 普通文本片段如果只用findall匹配 special token pattern它只会返回匹配到的 special token而不会返回中间那些普通文本。这个点当时很容易误解因为findall看起来像“正则化切分”但它本质上是“找出所有匹配项”不是“按规则切开并保留剩余文本”。另外special tokens 需要按长度从长到短排序。原因是可能有重叠 special token|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是pair - set[token_tuple]因为训练阶段的数据结构是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 里的反向索引更适合写成pair - set[token_index]也就是记录某个 pair 出现在哪些token_list下标里。这样每轮 merge 时只处理受影响的 token找到包含 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 里的核心差异train_bpe: pair - token_tuple 目标是更新 Counter、stats、pair_to_token encode: pair - token_list index 目标是更新当前输入文本的 token_list2.6 最后一步 flatten不能把整个list[bytes]直接查 vocabencode 的最后输出必须是list[int]但内部的token_list是混合结构int list[bytes]所以最后 flatten 时要分情况如果 token 是 int说明它是 special token id直接 append 如果 token 是 list[bytes]就遍历里面每个 bytes再用 rev_vocab 查 id当时一个容易错的地方是想直接对整个token做反查rev_vocab[token]但此时token是一个列表比如[bHe, bllo]它不是单个 bytes也不能作为rev_vocab的 key。真正应该反查的是列表里的每一个 bytes tokenbHe - id bllo - id这一步把前面的内部 byte-level 表示最终转换回作业要求的 token id。2.7decode反而很简单先拼 bytes再统一 UTF-8 解码相比encodedecode简单很多。因为vocab本来就是id - bytes所以 decode 只需要根据 id 找到 bytes 把所有 bytes 拼起来 最后整体 decode 成 str这里重要的是“整体 decode”而不是每个 token 单独 decode。原因是 Unicode 字符可能由多个 bytes 组成如果每个 token 单独 decode可能会把一个字符拆坏。所以正确思路是ids - bytes list - b.join(...) - decode(utf-8, errorsreplace)errorsreplace的作用是遇到非法 UTF-8 字节序列时不直接抛异常而是用替换字符处理。这和 tokenizer 测试里的鲁棒性要求更匹配。3. 流式编码encode_iterable3.1 真正麻烦的是流式边界最后是encode_iterable。这个函数的接口是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 的核心仍然是统计 pair 按 merges 顺序合并 输出 token id真正复杂的是把这个核心逻辑放进 CS336 的接口和测试语义里。第一train_bpe并行化时要先想清楚进程之间传什么。直接传大量 pre-token 列表能跑但不够紧凑让子进程直接返回Counter主进程再合并是更自然的拆分。这个过程也顺便弄清楚了AsyncResult.get()、Pool生命周期、Counter聚合和 Windows 多进程入口保护。第二encode的核心问题是类型适配。输入是str但 vocab 和 merges 都是 byte-levelvocab是int - bytesmerges是tuple[bytes, bytes]。所以普通文本必须转成list[bytes]后才能参与 merge最后再通过rev_vocab转成 id。中间只要混进 int、str、list[bytes] 的错误使用merge 就可能完全匹配不上。第三special token 必须在最开始识别。它不是普通文本的一部分而是一个整体 token。如果先拆成字节再试图在最后识别 special token边界信息已经丢了后续 merge 也可能改变它的内部结构。第四pair_to_token的思路可以从train_bpe迁移到encode但不能照搬。训练阶段存的是不可变 token tuple因为要维护Counter、stats和pair_to_tokenencode 阶段更适合存 token 下标因为当前输入的token_list会被原地更新。第五decode很直接但也有一个原则先把所有 id 对应的 bytes 拼起来再整体 UTF-8 decode。不要逐 token 解码否则可能拆坏多字节 Unicode 字符。第六encode_iterable暴露的是流式边界问题。严格通用方案需要维护 buffer、识别 stable prefix、保留 unsafe suffix但在当前测试输入按文件行迭代的前提下直接逐段encode再yield from是一个合理的工程取舍。所以这次实现最终可以概括成一句话train_bpe 解决的是如何训练出符合 CS336 要求的 byte-level vocab 和 merges Tokenizer 解决的是如何把 str 输入、安全的 special token 处理、byte-level merge 规则和最终 token id 输出接成一条完整链路。classTokenizer:def__init__(self,vocab,merge_dic,special_tokensNone):self.vocabvocab self.rev_vocab{v:kfork,vinvocab.items()}# 翻转 vocab ,方便后续 encode 从 bytes 找到 token_idself.merge_dicmerge_dic self.special_tokensspecial_tokensor[]ifself.special_tokens:self.special_tokens.sort(keylen,reverseTrue)self.set_special_tokensset(self.special_tokens)new_st[]forstinself.special_tokens:stre.escape(st)new_st.append(st)self.split_parten((|.join(new_st)))ifnew_stelse[]self.gpt2patre.compile(rs|t|re|ve|m|ll|d| ?\p{L}| ?\p{N}| ?[^\s\p{L}\p{N}]|\s(?!\S)|\s)classmethoddeffrom_files(cls,vocab_filepath,merges_filepath,special_tokensNone):withopen(vocab_filepath,r,encodingutf-8)asf:vocab0json.load(f)vocab{token:bytes(bytes_list)fortoken,bytes_listinenumerate(vocab0)}withopen(merges_filepath,r,encodingutf-8)asf:merge0json.load(f)merge[(bytes(left),bytes(right))forleft,rightinmerge0]returncls(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]] ] defencode(self,text):ifself.split_parten:text_listre.split(self.split_parten,text)new_text_list[]fortextintext_list:iftextinself.set_special_tokens:new_text_list.append(self.rev_vocab[text.encode(utf-8)])continuesub_text_list[text0.encode(utf-8)fortext0inre.findall(self.gpt2pat,text)]new_text_list.extend(sub_text_list)else:new_text_list[text0.encode(utf-8)fortext0inre.findall(self.gpt2pat,text)]text_listnew_text_list# 正则化token_list[]fortextintext_list:ifisinstance(text,int):token_list.append(text)continuenew_token[]fortext0intext:new_token.append(bytes([text0]))token_list.append(new_token)# 把特殊字符串直接转成 token_idpair_to_tokendefaultdict(set)foridx,tokeninenumerate(token_list):ifisinstance(token,int):continueforpairinzip(token,token[1:]):pair_to_token[pair].add(idx)formerged_pairinself.merge_dic:old_idx1merged_pair[0]old_idx2merged_pair[1]new_bytesb.join(merged_pair)ifmerged_pairnotinpair_to_token:continueforaffected_token_idxinlist(pair_to_token[merged_pair]):tokentoken_list[affected_token_idx]new_token[]i0while(ilen(token)):if(ilen(token)-1andtoken[i]old_idx1andtoken[i1]old_idx2):new_token.append(new_bytes)i2else:new_token.append(token[i])i1token_list[affected_token_idx]new_tokenforpairinzip(token,token[1:]):ifpairnotinpair_to_token:continuepair_to_token[pair].discard(affected_token_idx)ifnotpair_to_token[pair]:delpair_to_token[pair]forpairinzip(new_token,new_token[1:]):pair_to_token[pair].add(affected_token_idx)fin_token[]fortokenintoken_list:ifisinstance(token,int):fin_token.append(token)continueforbyteintoken:fin_token.append(self.rev_vocab[byte])returnfin_tokendefdecode(self,ids):text_bytesb.join(self.vocab[idx]foridxinids)returntext_bytes.decode(utf-8,errorsreplace)defencode_iterable(self,iterable):fortextiniterable:tokenself.encode(text)yieldfromtoken