跳转到内容

基础原理第 8 篇,共 13 篇

训练和推理:同一模型,两种序列组织

统一教学例子中的标签错位、训练张量、Prefill、Decode 与 KV 时序

更新于 阅读约 3 分钟

训练用正确序列监督每个位置的下一 token 分布。生成没有未来答案,只能先选一个 token,再用它预测后续 token。核心结构相同,输入组织与后续操作不同。

先读:模型流程、因果依赖。接着读:loss、缓存内部。

完整教学序列:小猫11 → 坐在23 → 窗边37 → 晒太阳42
input_ids: [[11, 23, 37]] [1,3]
labels: [[23, 37, 42]] [1,3]
Embedding: [1,3,8]
各层 Q/K/V: [1,2,3,4]
scores: [1,2,3,3],每行遮住未来 key
final hidden: [1,3,8]
logits: [1,3,100]

位置 0 读“小猫”,监督答案为“坐在”。位置 1 读“小猫 坐在”,监督答案为“窗边”。位置 2 读前三个 token,监督答案为“晒太阳”。Teacher forcing 指训练使用正确历史;不是把当前要预测的标签送给当前位置。

标签比输入左移一位。实际库可能接收完整 input_ids 和同长度 labels,在模型内部 shift;也可能由数据管线预先对齐。不要同时做两次 shift。教学写法明确输入/标签关系,并非所有库都按同样参数接收它们。

把 logits reshape 为 [3,100]、标签为 [3],可用 CrossEntropyLoss 计算三个位置的平均 loss。输入应为 logits,通常不先做 softmax。SFT 常把不监督的 prompt/padding 位置设为 ignore_index,例如 −100;应按实际模板和任务确定监督区域。

假设每步选中预设 token,不是实际模型运行结果:

forward本次输入Q使用的完整 K/V(每层各一份)scoresforward 后选出
Prefill[11,23][1,2,2,4][1,2,2,4],cache 长 2[1,2,2,2]37“窗边”
Decode 1[37][1,2,1,4][1,2,3,4],cache 长 3[1,2,1,3]42“晒太阳”
Decode 2[42][1,2,1,4][1,2,4,4],cache 长 4[1,2,1,4]下一个 token

Prefill 的 LM Head 可得到 [1,2,100],生成通常用最后一行。增量 Decode 得到 [1,1,100]。一些推理实现直接只对所需 hidden states 计算 LM Head,无需保存全部位置的 logits。

37 刚被选出来时只是一个 ID,cache 仍长 2。 下一轮输入 37,经过每层投影后,才把它的 K/V 追加到该层缓存,长度变 3。42 同理。位置索引也要继续递增,不能每轮从位置 0 开始。

问题训练常规生成
上文来源正确历史prompt + 模型此前输出
读取未来不允许,由 causal mask 保证未来不存在
使用哪个位置所有被监督的位置通常最后一个查询位置
梯度/参数更新求梯度,再由 optimizer 更新通常关闭梯度,不更新参数
KV通常不需要持久生成缓存复用每层历史 K/V

完整序列训练和 Prefill 通常适合大矩阵乘法;小 batch 单步 Decode 常受权重/KV 的内存读写影响。瓶颈依模型、硬件、batch 和上下文长度,不应写成永远 compute-bound / memory-bound。模型可以用同一类 Attention 算法或不同 kernel。

训练历史与生成历史的分布差异称 exposure bias。可理解这个现象,但不能把所有后训练方法都等同于直接修复它。

更完整的数字计算:D=2 手算。来源与代码:Karpathy GPT 训练/生成。