填充浪费计算,打包模糊文档边界
一句话总结
预训练数据会做成一条很长的“令牌序列”,切成批次喂给模型。MiniMind 把一篇文档放在一行,并 填充 到最大长度——简单,文档也不会混在一起,但文档一短,大部分计算就浪费在填充上。把多篇文档首尾相接、填满一行的 打包(packing) 没有被浪费的计算,代价是同一个窗口里,前一篇文档会被后一篇文档看到。这个模块会在我们的语料上用数字量出这两种方式的成本。
为什么需要它
模型接收的是形状为 (배치, 길이) 的整数张量。文档的长度各不相同,而张量必须是矩形,所以总得在某个地方对齐。办法有两种。
填充。一篇文档占一行,不足的位置用填充令牌补齐。MiniMind 的 PretrainDataset 就是这样做的。
tokens = tokenizer(text, max_length=self.max_length - 2, truncation=True).input_ids
tokens = [bos_token_id] + tokens + [eos_token_id]
input_ids = tokens + [pad_token_id] * (self.max_length - len(tokens))
labels = input_ids.clone(); labels[input_ids == pad_token_id] = -100
填充位置的标签是 -100,不会进入损失,但 计算照样会做。矩阵乘法不知道哪里是填充。我们语料里的文档平均 33 个令牌,如果把最大长度设为 128,超过 70% 的计算都花在了填充上。反过来,如果把最大长度设得短,长文档就会被截断。MiniMind 的 README 也写明了这个取舍——短样本因填充而浪费计算,长样本被截断而丢失信息。所以它按数据分别标出推荐的 max_seq_len。
打包。给每篇文档套上 [bos] … [eos],首尾相接连成一条长序列,再把这条序列按固定长度切开。没有被浪费的位置。代价是一个窗口里会放进三四篇文档,而因果掩码让每个令牌看到“前面所有的令牌”,所以后一篇文档的令牌会看到前一篇文档。大多数预训练都接受这一点——[eos] 和 [bos] 会告知边界,模型也会学会忽略边界之外的内容。想不让它越过边界,就得另外为每篇文档构造注意力掩码,实现会因此变得更复杂。
工作原理
打包的结果是一个整数数组。词表小于 65,536 时,uint16 就够了,只需要 int64 四分之一的空间。训练循环从这个数组里任选位置,截取 길이 那么长,组成批次。
ix = torch.randint(0, len(train) - seq, (batch,), generator=g)
x = torch.stack([train[i:i + seq] for i in ix])
loss = model(x, labels=x).loss # 한 칸 미는 일은 모델 안에서
MiniMind 模型把与输入相同的张量作为 labels 接收,在内部把 logits[..., :-1] 和 labels[..., 1:] 对齐,算出下一个令牌预测的损失。所以数据这一侧不需要另外把输入和答案错开一位。
验证数据必须 按文档 拆分。如果把一个数组按前 95%、后 5% 切开,一篇文档可能被劈成两半,横跨两边;如果同一篇文档既在训练里又在验证里,验证损失测的就是背下来的东西。拆分之后,还要确认一次同样的文本是否原样出现在训练一侧。
在现场相遇的样子
像简短问答、聊天记录这种文档很短的数据,如果用填充来训练,GPU 忙个不停,损失却下降得很慢。利用率很高,但真正在学习的令牌很少。这时只要数一数每一步 真正的令牌数,原因就一目了然。反过来,换成打包之后,如果模型开始越过文档边界,胡乱接续不相干的内容,就先检查 [eos] 是否贴对了。
训练预算也是按令牌来算的。“多少步”的含义会随批次大小和长度而变,所以要先算好每一步吃掉多少令牌,以及用语料的全部令牌走一遍(一个轮次)需要多少步,才能读懂损失曲线的起伏。
本课程与 MiniMind 原版的不同之处
MiniMind 用 HuggingFace datasets 读取 jsonl,在每个样本的 __getitem__ 中当场运行分词器。这样不必事先把 1.2GB 的语料转成令牌,代价是训练期间 CPU 要不停地切分令牌(所以默认的 num_workers 是 8)。这门课程的语料不到 2MB,所以一次性全部切好,存成一个 uint16 数组,训练循环只需从这个数组里截取窗口来用。在大规模预训练中,这种方式(事先令牌化并存成二进制文件)也很常见——因为反复遍历数据时,不必重复同样的令牌化,而且可以对文件做内存映射,只读取需要的部分。
下一项实验要做什么
用基准分词器测量文档长度,数出 MiniMind 式填充浪费了多少。把训练和验证语料按文档打包,存成 uint16 数组,并计算验证文档是否泄漏到了训练一侧,打包后的窗口里有百分之几的令牌能看到前一篇文档,以及走一遍需要多少步。