Skip to content

推理:从概率分布到逐 token 生成 ​

训练一次前向可同时评估所有位置,生成却必须等当前 token 被选出才能构造下一步输入。这是自回归模型在训练与服务端呈现不同性能形态的根源。理解 sampling、KV cache 与位置编号,能让你区分“输出不好”和“程序计算错了”。

为什么只取最后一个位置 ​

给定 prompt 的 token x1:t,模型输出 [B,t,V] 的 logits。最后一行表示 p(xt+1∣x1:t)。从这里选出一个 token,追加到序列,重复即可。前面各行预测的是 prompt 内已有的后继,不应拼成新的回复。课程实验的 generate 显式保留这条循环,便于断点检查。

python
for _ in range(new_tokens):
    window = ids[-model.config.context:]
    logits = model.forward(np.array([window]))[0][0, -1]
    next_id = choose(logits)
    ids.append(next_id)

上面 choose 表示采样策略;本地实现直接写在函数内。空 prompt 没有最后位置,因此要求至少一个 token;有 BOS 的真实系统可以用 BOS 开始。我们的 corpus 未训练 EOS,生成按 token 上限停止,不应把换行自动解释为生成完成。

temperature、top-k 与 top-p ​

温度 τ>0 时使用 softmax(z/τ)。小温度放大分数差异,分布更尖;大温度使其更平。τ=0 不应真的作除法,而应明确进入 argmax 分支。贪心选择每步局部最大概率,并不保证整个序列概率最大。

top-k 只保留分数最高的 k 项,再归一化。top-p 则排序概率,保留累计概率首次达到阈值 p 的最小前缀,并保留越过阈值的那一项。两者筛掉的候选不同;如果组合,顺序会影响结果。本地实现只有 temperature 与 top-k,没有伪装提供 top-p。改变随机 seed 能生成不同样本,但比较模型时应固定采样设置,否则质量差异可能来自解码器。

一个小例子:概率为 [0.55,0.25,0.15,0.05],top-k=2 留前两项;top-p=0.9 要留前三项,因为前两项之和只有 0.8。筛选后必须重新归一化。词表掩码也应用于分数或概率后再归一化,不能简单把某个 ID 替换成另一个 ID。

KV cache 为什么成立 ​

对因果模型,加入未来 token 不会改变已算出的历史 key/value。于是每层保存 K1:t,V1:t,下一步只算新位置的 query/key/value,用一个 query 读取全部历史。没有 cache 时反复计算整个前缀;有 cache 时,投影和 FFN 的历史计算可以省掉。每步的注意力仍需读取越来越长的 cache,并未变成常数成本。

普通多头注意力的 cache 内存约为

2BLTHdhs=2BLTdsbytes,

其中 s 是每元素字节数,2 表示 K 和 V。示例 B=1,L=2,T=128,d=64,s=2,cache 为 65536 bytes,即 64 KiB;不含模型权重与工作区。GQA/MQA 减少 KV 头数,公式相应把 H 换为 Hkv。

flowchart LR
  P[Prompt prefill] --> C[各层历史 K V]
  X[最新 token] --> Q[新 Q K V]
  C --> A[新 Q 读取缓存]
  Q --> A
  A --> O[下一 token 分布]
  Q --> U[追加新 K V]
  U --> C
查看流程图文本
flowchart LR
  P[Prompt prefill] --> C[各层历史 K V]
  X[最新 token] --> Q[新 Q K V]
  C --> A[新 Q 读取缓存]
  Q --> A
  A --> O[下一 token 分布]
  Q --> U[追加新 K V]
  U --> C

prefill 有大量 token 的矩阵计算,decode 每次只处理少量新 token,常更受内存搬运与批处理效率制约。吞吐、首 token 延迟和每 token 延迟不是同一个指标,服务优化需要同时报告。

RoPE 与滑动窗口的陷阱 ​

cached key 已包含位置旋转;新 query 必须使用正确的绝对位置或一致的相对偏移。不能每一步都把新 token 当成位置零而沿用旧 cache。截掉左侧 cache 后,既可以保留原位置编号,也可以对所有保留项一致调整,但不能只改 query。

本地 CPU 实验没有实现 KV cache:每步截取最近 context 个 token 并重新前向,窗口内位置从零开始。对 RoPE 的同一窗口做共同平移不改变其相对角度,但这仍是有限上下文重算策略,与保留长历史状态的缓存解码不同。它的目的在于验证训练与生成接口,不提供服务性能结论。读者实现 cache 后,应在不截断的序列上比较每步 logits 与全前缀重算,误差应处于数值精度范围。

自测与练习 ​

还应区分“停止生成”与“删掉已生成内容”。若使用 EOS,先判断新 token 是否等于 EOS,再按展示策略决定是否把标记显示出来;不能仅凭输出字符串末尾出现某个普通词就停止。batch 中不同请求可能在不同时间完成,已完成请求不应继续采样或计入延迟。测试 cache 时也要覆盖一个 token 的 prompt、长度刚好达到上限、超过上限后的窗口策略,以及不同请求的 cache 不被串用。这些边界条件往往比长段随机文本更容易暴露索引错误。

分别使用 temperature=0、0.7、1.3,从同一 checkpoint 和 prompt 生成。记录重复率与非法字节显示,不把三个样本直接当作统计结论。默认字符词表只支持训练中见过的字符;BPE 支持任意 UTF-8 输入,但语料没教过的中文不会因此自动获得语言能力。

推理设置 eval() 就不需要 no_grad() 了吗?

不够。eval() 修改 dropout、BatchNorm 等层的行为,no_grad() 或 inference mode 则决定是否记录自动微分状态。服务时通常两者都需要。本地手写 NumPy 前向仍保存了缓存,属于可读性优先的额外开销。

来源:2025 第 10 讲入口、PagedAttention/vLLM 论文。下一章:系统优化。