一句话总结:训练时,正确的前文和目标都已经写在数据里,所以同一条序列的所有位置可以一起计算;推理时,下一个 Token 还不存在,必须先选出它,才能继续下一步。串行的是一条序列里的生成步骤,不是所有请求都只能排队。
16.1 训练 vs 推理:核心区别
16.1.1 先看一张表
| 训练(Training) | 推理(Inference) | |
|---|---|---|
| 目的 | 让参数学会预测下一个 Token | 用已经学好的参数生成内容 |
| 输入 | 一段已知文本,输入与目标错开一位 | Prompt 加上已经生成的 Token |
| 目标 | 每个位置的正确下一个 Token 都已知 | 下一个 Token 未知,要由解码策略选择 |
| 序列内计算 | 一次前向传播算出所有位置的 logits | Prompt 可一次处理;之后的生成步骤前后依赖 |
| 梯度与参数 | 算梯度,并由优化器更新参数 | 不算梯度,也不更新参数 |
这里的“并行”和“串行”,说的是同一条序列里的位置:
- 训练仍然要做反向传播,整体成本当然不会比一次推理便宜。
- 推理系统可以把很多请求组成 Batch,多个请求同时跑。
- 但对某一条回答来说,第 2 个新 Token 必须等第 1 个选出来以后才有完整前文。
16.1.2 为什么会有这个区别?
还是沿用原书的例子:
- 训练文本已经写完整:“小沈阳江西演唱会邀请了沈春阳”
- 推理时手里只有:“小沈阳江西演唱会邀请了”
- 训练数据能直接告诉模型每个位置后面是什么;推理时没有这份答案
注意,这里展示的是可读文字,不是在断言“一个汉字等于一个 Token”。实际切分由 Tokenizer 决定;同一个汉字甚至可能被拆成多个 byte-level Token。第 4 章已经见过这种情况。
16.2 训练过程详解
16.2.1 Teacher Forcing 到底是什么?
假设一段文本被 Tokenizer 切成:
完整序列:[x0, x1, x2, ..., xT]
模型输入:[x0, x1, x2, ..., xT-1]
训练目标:[x1, x2, x3, ..., xT]
目标并不是含糊地“右移一下”,而是比同列输入早走一步:
- 位置 0 读到 x0,预测 x1
- 位置 1 读到 x0、x1,预测 x2
- ...
- 位置 T-1 读到 x0 到 xT-1,预测 xT
训练时喂给模型的是数据里的正确前文,而不是模型上一位置自己猜出来的结果。这个做法通常叫 Teacher Forcing;对 Decoder-only 语言模型来说,它就是标准的 next-token training。
16.2.2 为什么所有位置能一起算?
因为 x0 到 xT 都已经在训练样本里,GPU 可以用一个带 Causal Mask 的矩阵运算,同时得到 T 个位置的 logits:
import torch.nn.functional as F
def train_step(model, optimizer, tokens):
"""tokens: [batch, T + 1];这里先省略 Padding。"""
model.train()
optimizer.zero_grad(set_to_none=True)
input_ids = tokens[:, :-1]
target_ids = tokens[:, 1:]
logits = model(input_ids) # [batch, T, vocab_size]
loss = F.cross_entropy(
logits.reshape(-1, logits.size(-1)),
target_ids.reshape(-1),
)
loss.backward() # 只计算梯度
optimizer.step() # 到这里才更新参数
return loss.detach()
如果 Batch 里有 Padding,还要同时做两件事:Attention 不能读 Padding;Padding 对应的 target 也不能进入 loss(PyTorch 里常把这些 target 改成 -100)。
16.2.3 Causal Mask 遮住的是什么?
第 t 个位置可以看到 x0 到 xt,并用它们预测 x(t+1):
位置 0 看到:[x0, -, -, -, ...] → 预测 x1
位置 1 看到:[x0, x1, -, -, ...] → 预测 x2
位置 2 看到:[x0, x1, x2, -, ...] → 预测 x3
对角线上的当前 Token 是可以看的;被遮住的是它后面的 Token。没有这个 Mask,位置 t 就可能直接读到自己的答案 x(t+1),训练变成抄答案。
16.3 推理过程详解
16.3.1 Prefill 和 Decode
标准的自回归推理分两段:
- Prefill:Prompt 已经完整存在,可以一次处理所有 Prompt Token。
- Decode:从最后位置的 logits 选一个新 Token,把它接回上下文,再预测下一个。
Prompt:“小沈阳江西演唱会邀请了”
↓ Prefill
选择下一个 Token(把结果解码成可读文字后可能是“沈”)
↓ 接回上下文
选择再下一个 Token(可能让文字继续成“沈春”)
↓
继续,直到 EOS、stop 条件或 max_new_tokens
上面写“沈”“春”只是为了让人读懂。程序真正追加的是 Token ID;一次 Token 可能对应一个字、半个字节片段、一个词的一部分,甚至带着前导空格。
16.3.2 KV Cache 改变了什么?
最朴素的代码会在每一步重新计算整个可见窗口。KV Cache 则保存过去各层的 K、V:
- Prefill 时把 Prompt 的 K、V 算好并存起来
- Decode 时通常只把最新 Token 送进各层,再读取缓存
- 新 Token 仍然要按顺序选择;省掉的是重复计算,不是前后依赖
第 22 章会把这件事拆开讲。
2026 补充:标准解码的逻辑仍然是有序的,但 Speculative Decoding 可以先用小模型草拟多个 Token,再让大模型并行验证。即使第一个草稿就被拒绝,校正步骤也会按目标分布产出一个 Token;如果草稿吻合,一轮则可能前进多个 Token。它减少大模型调用轮数,并没有把目标模型变成非自回归模型。
16.3.3 Context Length 不是自动滑动开关
Prompt Token 加已生成 Token 不能超过模型支持的上下文。超过以后,应用程序必须决定怎么办:
- 停止生成或拒绝过长输入
- 截掉较早的 Token
- 先摘要,再把摘要放回上下文
- 使用模型原生支持的 Sliding-window Attention
“永远保留最近 N 个 Token”只是其中一种策略,不是所有 GPT 都会自动替你做。Token 一旦被截掉,就不再是本次前向传播的输入;这比说模型“心理上忘了”更准确。
16.4 一个不带 KV Cache 的基线循环
下面的代码故意保留最朴素的写法,好让串行依赖一眼可见。它只支持 Batch size = 1,而且每一步都会重新计算当前窗口;不要把它误认为生产推理器。
import torch
@torch.inference_mode()
def generate_greedy(
model,
prompt_ids,
context_length,
eos_token_id,
max_new_tokens=50,
):
"""prompt_ids: [1, prompt_length]"""
if prompt_ids.ndim != 2 or prompt_ids.size(0) != 1:
raise ValueError("This teaching loop expects batch size 1")
model.eval()
generated = prompt_ids.clone()
for _ in range(max_new_tokens):
model_input = generated[:, -context_length:]
logits = model(model_input)
next_token = logits[:, -1, :].argmax(dim=-1, keepdim=True)
generated = torch.cat((generated, next_token), dim=1)
if next_token.item() == eos_token_id:
break
return generated
几个容易混淆的点:
- model.eval() 只切换 Dropout 等模块的行为,不会自动选择“最可能的词”。
- torch.inference_mode() 关闭梯度跟踪,进一步省掉 Autograd 开销。
- 生成下一个 Token 时,只需要最后位置的 logits。
- argmax 是 Greedy Decoding;换成采样后,即使 model.eval(),输出仍然可以不同。
16.5 Padding 和批量推理
16.5.1 左 Padding 为什么方便?
不同 Prompt 的 Token 数不同,组成矩阵前必须补到一样长。对 Decoder-only 推理,常见做法是左 Padding:
请求 A:[a0, a1, a2, a3] mask = [1, 1, 1, 1]
请求 B:[PAD, PAD, b0, b1] mask = [0, 0, 1, 1]
这样最后一列对每个请求都是真实 Token,取 logits[:, -1, :] 比较直接。上面的 a、b 才是 Token;它们可以分别来自“小沈阳江西演唱会邀请了”和“请续写”这样的真实文字。
16.5.2 两种 Mask 不要混在一起
- Causal Mask:不许看未来位置。
- Padding Mask:不许把 PAD 当成上下文。
实现通常会把两者合并,在 Softmax 之前把被遮住的 Key 分数变成负无穷,因此它们的 Attention 权重为 0。模型在推理时没有“学会忽略”Padding——是 Mask 当场强制它忽略。
tokenizer.padding_side = "left"
batch = tokenizer(prompts, padding=True, return_tensors="pt")
logits = model(
batch["input_ids"],
attention_mask=batch["attention_mask"],
)
next_token_logits = logits[:, -1, :]
如果使用右 Padding,最后一列可能是 PAD,就不能对所有请求盲目取 -1;还要找到每行最后一个真实位置。使用 learned absolute position embedding 时,自写左 Padding 代码还必须正确处理 position IDs。成熟框架的 generate 通常会替你处理,但自制模型不能靠猜。
批量推理、Continuous Batching 可以让不同请求一起跑;它们提高设备利用率,却没有消除每条回答内部的 Token 依赖。
16.6 计算与运行模式的真实对比
| 训练 | 自回归推理 | |
|---|---|---|
| 一次样本的主流程 | 一次前向 + 一次反向 + 优化器更新 | 一次 Prefill + N 个相互依赖的 Decode 步骤 |
| 可并行的维度 | Batch、序列位置、Attention head | Batch、Prompt 位置;不同请求也可并行 |
| 必须保存的状态 | 激活、梯度,通常还有优化器状态 | 权重,以及常见的 KV Cache |
| 常见瓶颈 | 算力、显存容量、通信 | Prefill 偏算力;Decode 常受显存带宽与单步延迟限制 |
| 参数是否改变 | optimizer.step() 时改变 | 不改变 |
因此,不要简单记成“推理比训练慢”。训练一整个模型的总计算量远大于生成一次回答;这里说的慢,是单条回答的新增 Token 无法像训练目标位置那样一次全部算完。
16.6.1 Dropout 的行为
如果模型里配置了 Dropout:
model.train() # Dropout 参与随机丢弃
model.eval() # Dropout 关闭
有些现代 LLM 本来就把 Dropout 设为 0。即使关闭 Dropout,只要解码使用 Sampling,输出仍然不是确定的;反过来,Greedy 也可能受底层数值非确定性影响。eval() 的准确含义是“切换到评估行为”,不是“保证每次文字相同”。
16.7 解码策略
模型给出的是 logits。怎么从中选下一个 Token,是解码器的工作。
16.7.1 Greedy 和直接采样
# Greedy 不需要先算 Softmax;argmax(logits) 与 argmax(Softmax(logits)) 相同
next_token = logits.argmax(dim=-1, keepdim=True)
# Sampling 要先得到概率;temperature 必须大于 0
probs = F.softmax(logits / temperature, dim=-1)
next_token = torch.multinomial(probs, num_samples=1)
- Greedy:每一步选最大 logit,容易复现,但局部最优不等于整段最优,也可能重复。
- Sampling:按分布抽样,结果更多样,也可能抽到低质量 Token。
- Temperature:T 小于 1 时分布更尖,T 大于 1 时更平;“T = 0”应实现成 Greedy,不能真的拿 logits 除以 0。
16.7.2 正确的 Top-K
top_k_probs, top_k_ids = torch.topk(probs, k=50, dim=-1)
sampled_slot = torch.multinomial(top_k_probs, num_samples=1)
next_token = top_k_ids.gather(-1, sampled_slot)
原来的索引写法在有 Batch 维度时会取错轴;gather 才是在每一行里把“第几个候选”还原成词表 Token ID。
16.7.3 正确的 Top-P(Nucleus Sampling)
def sample_top_p(logits, top_p=0.9, temperature=1.0):
if not 0.0 < top_p <= 1.0:
raise ValueError("top_p must be in (0, 1]")
if temperature <= 0:
raise ValueError("temperature must be positive")
sorted_logits, sorted_ids = torch.sort(
logits / temperature, descending=True, dim=-1
)
sorted_probs = F.softmax(sorted_logits, dim=-1)
cumulative = torch.cumsum(sorted_probs, dim=-1)
# 保留“让累计概率第一次达到 top_p”的那个 Token
remove = cumulative > top_p
remove[..., 1:] = remove[..., :-1].clone()
remove[..., 0] = False
filtered_logits = sorted_logits.masked_fill(remove, float("-inf"))
filtered_probs = F.softmax(filtered_logits, dim=-1)
sampled_slot = torch.multinomial(filtered_probs, num_samples=1)
return sorted_ids.gather(-1, sampled_slot)
原来的 cumsum <= 0.9 会把刚好跨过阈值的那个 Token 排除,极端时甚至一个候选都不留。Top-P 的定义是保留累计概率达到 P 的最小集合。实际系统常把 Temperature、Top-K、Top-P 组合使用;它们不是互斥开关。
16.8 为什么主流 GPT 使用自回归?
16.8.1 真正的理由不是“语言有顺序”这么简单
“我爱你”和“你爱我”当然不同,但能表示顺序的模型不只有自回归模型。自回归建模选择的是下面这种概率分解:
p(x1, x2, ..., xT) = Π p(xt | x1, ..., x(t-1))
它的好处是:
- 每一步都有明确的 next-token 训练目标
- 可以表示可变长度序列
- 生成时能把 Prompt 和已生成内容都作为条件
- 似然可以按 Token 分解、计算和优化
但它不保证文本连贯,更不保证事实正确。它只是允许每一步使用此前上下文;最终质量还取决于模型、数据、解码和上下文。
16.8.2 非自回归并没有消失
非自回归机器翻译、迭代 Mask 预测和 Diffusion Language Model 都在探索并行生成。2025 年的 LLaDA 就展示了从头训练的大型 Diffusion Language Model。
所以,“非自回归一定更差”“最好的模型全是自回归”都说得太满。更准确的结论是:自回归 Decoder 仍是成熟而常见的路线;并行或迭代生成是有真实结果、也有不同速度与质量取舍的替代路线。
16.9 本章总结
16.9.1 核心对比
| 方面 | 训练 | 推理 |
|---|---|---|
| 目标已知? | 是,来自训练文本 | 否,要现场选择 |
| 序列内处理 | 所有目标位置一起算 | Prefill 后按依赖顺序 Decode |
| 计算阶段 | Forward + Backward + Optimizer | Forward only,通常带 KV Cache |
| Dropout | 按配置启用 | eval() 时关闭 |
| 参数更新 | optimizer.step() 时更新 | 不更新 |
16.9.2 核心认知
训练并行,是因为每个位置的正确前文已经在数据中;标准推理解码串行,是因为下一个 Token 会成为再下一个 Token 的前文。KV Cache 减少重复计算,Batching 并行不同请求,Speculative Decoding 减少大模型调用轮数,但三者都没有抹掉这条条件依赖链。
本章交付物
学完这一章,你应该能够:
- 用 x0 到 xT 写出 input 与 target 的一位错位
- 解释 Causal Mask 为什么允许看当前 Token、却不能看下一个 Token
- 区分 Prefill、Decode 和 KV Cache
- 说明为什么“单序列 Decode 串行”不等于“所有推理都不能并行”
- 写出不会漏掉阈值 Token 的 Top-P 采样
下一章预告
推理效率的主线已经看清了,KV Cache 留到第 22 章再拆。在那之前,还要补上训练里最敏感的一只旋钮——学习率。它太大时会让更新越过合适区域,太小时又走得太慢。下一章,我们来看它究竟在控制什么。