训练 vs 推理:同一个 Transformer,两条完全不同的执行路径
📅 创建时间:2026-07-29 🏷️ 标签:#Training #Inference #TeacherForcing #AutoRegressive #KVCache #FLOPs #Memory 📚 前置知识:[[08-transformer-by-hand]](理解 Self-Attention、QKV 和 Decoder 架构)
📋 本章目标
- 理解为什么训练和推理使用同一套模型参数,但执行路径截然不同
- 掌握 Teacher Forcing 的核心思想——"用正确答案当输入",以及 Causal Mask 如何让并行训练成为可能
- 掌握 Auto-Regressive 生成的完整流程——每次只吐一个 token,N 个 token 需要 N 次前向传播
- 理解 KV Cache 为什么是推理优化的灵魂——把 O(n²) 降到 O(n)
- 理解训练和推理的计算量差异(为什么训练一次前向 ≈ 推理 3x)
- 理解训练和推理的显存差异(为什么训练 7B 需要 ~100GB,推理只需 ~14GB)
- 理解为什么训练和推理的优化方向完全不同——一个拼吞吐,一个拼延迟
第0部分:同一个模型,两个完全不同的运行模式
0.1 一个很多人忽略的事实
当你加载 Llama-3-8B 的权重文件时,你加载的是一个训练完成的模型。这个模型在训练时见过几万亿个 token,经历过数百万次梯度更新。但当你用它来聊天时,它的参数是冻结的——没有任何梯度计算,没有任何权重更新。
训练和推理执行的是完全相同的 Transformer 计算(Self-Attention + FFN),但是:
- 训练时:喂入完整的目标序列,所有位置并行计算,计算损失,反向传播,更新权重。
- 推理时:只喂入起始标记,一个 token 一个 token 地串行生成,没有反向传播,权重不变。
同一个模型,两套完全不同的运行逻辑。
┌─────────────────────────────────────────────────────────────┐
│ 同一个 Transformer 模型,两种运行模式 │
├─────────────────────────────────────────────────────────────┤
│ │
│ ┌─────────────────┐ │
│ │ │ │
│ │ Transformer │ ← 同一套参数(权重矩阵 Wq,Wk,Wv,...) │
│ │ Model │ │
│ │ │ │
│ └────────┬────────┘ │
│ │ │
│ ┌───────┴───────┐ │
│ │ │ │
│ ▼ ▼ │
│ ┌──────────┐ ┌──────────┐ │
│ │ TRAINING │ │INFERENCE │ │
│ │ MODE │ │ MODE │ │
│ ├──────────┤ ├──────────┤ │
│ │ │ │ │ │
│ │ Forward │ │ Forward │ │
│ │ + │ │ ONLY │ │
│ │ Backward │ │ │ │
│ │ + │ │ (no grad)│ │
│ │ Update │ │ │ │
│ │ │ │ │ │
│ │ 输入: │ │ 输入: │ │
│ │ 完整序列 │ │ 逐 token │ │
│ │ 并行计算 │ │ 串行生成 │ │
│ │ │ │ │ │
│ │ 目标: │ │ 目标: │ │
│ │ 最小化 │ │ 生成 │ │
│ │ Loss │ │ 合理文本 │ │
│ │ │ │ │ │
│ │ 产出: │ │ 产出: │ │
│ │ 更新后 │ │ 一个 │ │
│ │ 的权重 │ │ token │ │
│ │ │ │ 序列 │ │
│ └──────────┘ └──────────┘ │
│ │
│ 核心矛盾:训练要并行(快),推理只能串行(无奈) │
│ 解决方案:Teacher Forcing 让训练"假装"并行 │
│ │
└─────────────────────────────────────────────────────────────┘0.2 为什么会有这种差异
根本原因只有一个:训练时有"正确答案",推理时没有。
训练时,你手里有完整的语料——"1+1=2"这个句子从头到尾都已经写好了。你可以把整个句子喂给模型,让模型在每个位置预测下一个 token,然后用正确答案计算误差。
推理时,你只有一个起始标记 <s>。模型必须先预测出第一个词,然后把这个词拼回去再预测第二个词……如果模型预测错了第一个词,后面的所有预测都会建立在错误的基础上。没有"正确答案"可以给它"兜底"。
第1部分:Training Mode(训练模式)——Teacher Forcing 深度剖析
1.1 Teacher Forcing 的核心思想
"Teacher Forcing" 这个名字很形象:
- Teacher(老师):手里有标准答案
- Forcing(强制):不管学生上一步预测了什么,老师都强制把正确答案作为下一步的输入
┌─────────────────────────────────────────────────────────────┐
│ Teacher Forcing 的直观理解 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 好比教小孩做加法: │
│ │
│ 学生要学:1 + 1 = 2 │
│ │
│ 普通教法(Auto-Regressive): │
│ 老师:"1" │
│ 学生:"+" ← 对了 │
│ 老师:"1 +" │
│ 学生:"=" ← 对了 │
│ 老师:"1 + 1 =" │
│ 学生:"3" ← 错了!接下来全错! │
│ │
│ Teacher Forcing: │
│ 老师不管学生预测了什么,每次都告诉学生正确答案: │
│ 位置0 输入 [<s>] → 学生猜 "1" → 正确答案是 "1" │
│ 位置1 输入 [<s>, 1] → 学生猜 "+" → 正确答案是 "+" │
│ 位置2 输入 [<s>,1,+] → 学生猜 "1" → 正确答案是 "1" │
│ 位置3 输入 [<s>,1,+,1] → 学生猜 "=" → 正确答案是 "=" │
│ 位置4 输入 [<s>,1,+,1,=] → 学生猜 "2" → 正确答案是 "2" │
│ │
│ 关键:输入永远是正确答案,不是学生自己的预测。 │
│ 这样错误不会累积——每个位置独立学习。 │
│ │
└─────────────────────────────────────────────────────────────┘1.2 具体例子:用 "1+1=2" 训练
假设我们用一个 Decoder-Only 模型(如 GPT)来学习 "1+1=2" 这个序列。
Tokenization(假设每个字符是一个 token):
词汇表:{<s>: 0, 1: 1, +: 2, =: 3, 2: 4, <eos>: 5}
目标序列:["<s>", "1", "+", "1", "=", "2", "<eos>"]
Token IDs:[0, 1, 2, 1, 3, 4, 5]准备训练数据——右移一位:
┌─────────────────────────────────────────────────────────────┐
│ 训练数据的构造:Input 右移 = Labels │
├─────────────────────────────────────────────────────────────┤
│ │
│ 原始序列: <s> 1 + 1 = 2 <eos> │
│ 训练输入: <s> 1 + 1 = 2 <eos> ← 去掉最后一个│
│ 训练标签: 1 + 1 = 2 <eos> ← 去掉第一个 │
│ │
│ 即: │
│ ┌──────────┬───┬───┬───┬───┬───┬──────┐ │
│ │ Position │ 0 │ 1 │ 2 │ 3 │ 4 │ 5 │ │
│ ├──────────┼───┼───┼───┼───┼───┼──────┤ │
│ │ Input │<s>│ 1 │ + │ 1 │ = │ 2 │ │
│ │ Label │ 1 │ + │ 1 │ = │ 2 │<eos> │ │
│ └──────────┴───┴───┴───┴───┴───┴──────┘ │
│ │
│ 模型的任务: │
│ 位置0,看到 [<s>] → 预测 "1" │
│ 位置1,看到 [<s>, 1] → 预测 "+" │
│ 位置2,看到 [<s>, 1, +] → 预测 "1" │
│ 位置3,看到 [<s>, 1, +, 1] → 预测 "=" │
│ 位置4,看到 [<s>, 1, +, 1, =] → 预测 "2" │
│ 位置5,看到 [<s>, 1, +, 1, =, 2] → 预测 <eos> │
│ │
└─────────────────────────────────────────────────────────────┘1.3 一次前向传播,所有位置并行计算
这是训练最神奇的地方:虽然模型在每个位置只能看到"当前位置及之前"的 token,但通过 Causal Mask,所有位置可以在一次矩阵乘法中并行计算。
┌─────────────────────────────────────────────────────────────┐
│ 一次 Forward Pass 同时计算所有 6 个位置 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 输入矩阵 X (6 × d_model): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ row 0: embedding of <s> │ │
│ │ row 1: embedding of 1 │ │
│ │ row 2: embedding of + │ │
│ │ row 3: embedding of 1 │ │
│ │ row 4: embedding of = │ │
│ │ row 5: embedding of 2 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Q = X × Wq (all positions at once) │
│ K = X × Wk (all positions at once) │
│ V = X × Wv (all positions at once) │
│ │
│ Scores = Q × K^T (6×6 matrix — all pairwise scores) │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ <s> 1 + 1 = 2 │ │
│ │ <s> [ s00, s01, s02, s03, s04, s05 ] │ │
│ │ 1 [ s10, s11, s12, s13, s14, s15 ] │ │
│ │ + [ s20, s21, s22, s23, s24, s25 ] │ │
│ │ 1 [ s30, s31, s32, s33, s34, s35 ] │ │
│ │ = [ s40, s41, s42, s43, s44, s45 ] │ │
│ │ 2 [ s50, s51, s52, s53, s54, s55 ] │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 然后施加 Causal Mask(下三角保留,上三角设为 -∞): │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ <s> 1 + 1 = 2 │ │
│ │ <s> [ s00, -∞, -∞, -∞, -∞, -∞ ] │ │
│ │ 1 [ s10, s11, -∞, -∞, -∞, -∞ ] │ │
│ │ + [ s20, s21, s22, -∞, -∞, -∞ ] │ │
│ │ 1 [ s30, s31, s32, s33, -∞, -∞ ] │ │
│ │ = [ s40, s41, s42, s43, s44, -∞ ] │ │
│ │ 2 [ s50, s51, s52, s53, s54, s55 ] │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Softmax 后,-∞ 位置的概率都是 0。 │
│ 位置 i 只能 attend 到位置 0,1,...,i。 │
│ │
│ Output = softmax(Masked_Scores) × V │
│ │
│ 一次矩阵乘法 → 同时得到 6 个位置的输出向量! │
│ │
└─────────────────────────────────────────────────────────────┘1.4 Causal Mask 的数学形式
Causal Mask 是一个上三角为 -∞ 的矩阵:
M[i][j] = 0 if j ≤ i (允许看到当前位置及之前)
M[i][j] = -∞ if j > i (禁止看到未来)
在代码中通常表示为:
┌ ┐
│ 0 -∞ -∞ -∞ -∞ -∞ │
│ 0 0 -∞ -∞ -∞ -∞ │
│ 0 0 0 -∞ -∞ -∞ │
│ 0 0 0 0 -∞ -∞ │
│ 0 0 0 0 0 -∞ │
│ 0 0 0 0 0 0 │
└ ┘
然后 Attention = softmax(QK^T / √d_k + M)
加 -∞ 的位置在 exp 后变成 0。1.5 计算 Loss——所有位置一起算
┌─────────────────────────────────────────────────────────────┐
│ 训练 Loss 的计算 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 前向传播后,每个位置输出一个 logits 向量 (1 × vocab_size): │
│ │
│ logits[0] = [0.1, 0.8, 0.05, 0.02, 0.02, 0.01] │
│ → 模型认为位置0最可能输出 token 1 ("1") │
│ │
│ logits[1] = [0.05, 0.05, 0.7, 0.1, 0.05, 0.05] │
│ → 模型认为位置1最可能输出 token 2 ("+") │
│ │
│ ...(共 6 个位置) │
│ │
│ Cross-Entropy Loss = -1/N × Σ log P(correct_token | logits)│
│ │
│ 对于位置0:correct_token = "1" (ID=1) │
│ loss[0] = -log(softmax(logits[0])[1]) │
│ = -log(0.36) = 1.02 │
│ │
│ 对于位置1:correct_token = "+" (ID=2) │
│ loss[1] = -log(softmax(logits[1])[2]) │
│ = -log(0.28) = 1.27 │
│ │
│ Total Loss = mean([loss[0], loss[1], ..., loss[5]]) │
│ │
│ 本质:让模型在每个位置都更可能输出正确答案。 │
│ 这是一个多分类问题(vocab_size 个类别)× N 个位置。 │
│ │
└─────────────────────────────────────────────────────────────┘1.6 反向传播——显存的真正杀手
┌─────────────────────────────────────────────────────────────┐
│ 训练循环的完整 4 步 │
├─────────────────────────────────────────────────────────────┤
│ │
│ Step 1: Forward Pass(前向传播) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Input → Embedding → [Transformer Block × L] │ │
│ │ → Linear → Softmax → logits │ │
│ │ │ │
│ │ 产出:每个位置的预测概率 │ │
│ │ 必须保存:所有中间激活值(activations) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 2: Loss Computation(计算损失) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Loss = CrossEntropy(logits, labels) │ │
│ │ │ │
│ │ 产出:一个标量值(比如 2.31) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 3: Backward Pass(反向传播) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ dL/dW = dL/d(output) × d(output)/d(W) │ │
│ │ │ │
│ │ 链式法则从输出端反向传播到输入端 │ │
│ │ → 需要 Step 1 中保存的所有激活值! │ │
│ │ → 计算量约 = 前向传播的 2 倍(矩阵乘法 + 转置乘法) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 4: Optimizer Update(权重更新) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ AdamW: │ │
│ │ m_t = β1 × m_{t-1} + (1-β1) × grad │ │
│ │ v_t = β2 × v_{t-1} + (1-β2) × grad² │ │
│ │ m̂_t = m_t / (1-β1^t) │ │
│ │ v̂_t = v_t / (1-β2^t) │ │
│ │ W_t = W_{t-1} - lr × m̂_t / (√v̂_t + ε) │ │
│ │ │ │
│ │ 需要额外存储:m(一阶动量)和 v(二阶动量) │ │
│ │ 每个参数存两份 → 参数量 × 8 bytes × 2 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 训练循环的 ASCII 图: │
│ │
│ ┌──────────┐ ┌──────────┐ │
│ │ Forward │────▶│ Loss │ │
│ │ Pass │ │Compute │ │
│ └──────────┘ └────┬─────┘ │
│ ▲ │ │
│ │ ▼ │
│ ┌────┴──────┐ ┌──────────┐ │
│ │ Update │◀────│ Backward │ │
│ │ Weights │ │ Pass │ │
│ └───────────┘ └──────────┘ │
│ │
│ 一次迭代 = Forward + Loss + Backward + Update │
│ 训练一个大模型需要数百万次这样的迭代。 │
│ │
└─────────────────────────────────────────────────────────────┘1.7 为什么训练时激活值必须保存
反向传播需要计算 dL/dW。以最简单的线性层 Y = X × W 为例:
前向:Y = X × W
反向:dL/dW = X^T × (dL/dY)
要计算 dL/dW,你需要:
- dL/dY(从上一层反向传回来,已知)
- X(前向传播时的输入,必须保存!)
如果你没保存 X,就得重新算一次前向传播——这就是 gradient checkpointing 的思想。对于 Attention 层,需要保存的激活值包括:
- Q、K、V(用于计算 dL/dWq, dL/dWk, dL/dWv)
- Softmax 之前的 Scores(用于计算 dL/dScores)
- Softmax 之后的 Attention Weights(用于计算 dL/dV)
所有这些矩阵的大小都是 batch_size × num_heads × seq_len × seq_len 量级。
第2部分:Inference Mode(推理模式)——自回归生成的完整旅程
2.1 推理的起点:只有一个 token
┌─────────────────────────────────────────────────────────────┐
│ 推理的完整流程:逐 token 生成 │
├─────────────────────────────────────────────────────────────┤
│ │
│ Step 1: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 输入:[<s>] │ │
│ │ 前向传播 → logits → Softmax → 概率分布 │ │
│ │ 从分布中采样 → token "I" │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 2: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 输入:[<s>, "I"] │ │
│ │ 前向传播 → logits → Softmax → 概率分布 │ │
│ │ 从分布中采样 → token "love" │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 3: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 输入:[<s>, "I", "love"] │ │
│ │ 前向传播 → logits → Softmax → 概率分布 │ │
│ │ 从分布中采样 → token "you" │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Step 4: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 输入:[<s>, "I", "love", "you"] │ │
│ │ 前向传播 → 预测 <eos> → 停止 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 4 个 token 的输出 = 4 次完整的前向传播。 │
│ 每次前向传播都要经过全部 L 层 Transformer Block。 │
│ │
│ 相比之下,训练只需要 1 次前向传播。 │
│ 这就是为什么"训练一个模型"很快,但"用模型生成文本"很慢。 │
│ │
└─────────────────────────────────────────────────────────────┘2.2 为什么推理不能像训练一样并行
一个常见的疑问:既然训练时可以通过 Causal Mask 让所有位置并行计算,推理时为什么不行?
答案:因为推理时你不知道后面的 token 是什么。
训练时:
Input = [<s>, 1, +, 1, =, 2] ← 你提前知道整个序列
第 3 个位置虽然只能看到前 3 个 token,
但第 3 个位置的 INPUT 是已知的(就是 "1")。
所以你可以把整个序列一次性输进去,用 Mask 限制可见范围。
推理时:
Input = [<s>] ← 你只知道第一个
你不知道第 1 个 token 是 "1" 还是 "+" 还是别的什么。
你必须先预测出 token 1,才能把 token 1 放到 input 里。
有了 token 1 之后,你才能预测 token 2。
...2.3 KV Cache——推理优化的灵魂
如果不做任何优化,推理时每一步都要重新计算所有 token 的 Attention。这意味着:
- Step 1:计算 1 个 token 的 Attention → O(1²) = O(1)
- Step 2:计算 2 个 token 的 Attention → O(2²) = O(4)
- Step 3:计算 3 个 token 的 Attention → O(3²) = O(9)
- Step N:计算 N 个 token 的 Attention → O(N²)
总计算量 = O(1 + 4 + 9 + ... + N²) = O(N³)
这是灾难性的。KV Cache 解决了这个问题。
┌─────────────────────────────────────────────────────────────┐
│ KV Cache 的原理 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 核心观察:每个 token 的 K 和 V 向量只依赖它自己的输入。 │
│ 新 token 不会改变已有 token 的 K 和 V。 │
│ │
│ 因此:计算过的 K 和 V 可以缓存起来! │
│ │
│ WITHOUT KV Cache:每一步重新计算所有 K 和 V │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Step 1: K₁,V₁, K₂,V₂, ..., Kₙ,Vₙ 全部重新算 │ │
│ │ Step 2: K₁,V₁, K₂,V₂, ..., Kₙ₊₁,Vₙ₊₁ 全部重新算 │ │
│ │ Step 3: ... │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ WITH KV Cache:只计算新 token 的 K 和 V │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Step 1: 计算 K₁,V₁ → 缓存到 cache[0] │ │
│ │ Step 2: 只算 K₂,V₂ → 追加到 cache[1] │ │
│ │ 读取 cache[0] 获得 K₁,V₁ │ │
│ │ Step 3: 只算 K₃,V₃ → 追加到 cache[2] │ │
│ │ 读取 cache[0:1] 获得 K₁:K₂, V₁:V₂ │ │
│ │ │ │
│ │ 第 N 步计算量:只算 1 个新 K 和 V │ │
│ │ Attention 计算:Q_new (1×d) × K_cache^T (N×d) │ │
│ │ → O(N) per step │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 总计算量从 O(N³) 降到 O(N²)! │
│ │
│ 对于生成 2048 个 token: │
│ 无 KV Cache:~14.3 billion operations │
│ 有 KV Cache:~2.1 million operations │
│ 差距近 7000 倍。 │
│ │
└─────────────────────────────────────────────────────────────┘2.4 KV Cache 的显存占用
KV Cache 不是免费的——它需要额外的显存:
┌─────────────────────────────────────────────────────────────┐
│ KV Cache 的显存计算 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 对于每一层 Transformer: │
│ K cache shape: (batch, num_heads, seq_len, d_k) │
│ V cache shape: (batch, num_heads, seq_len, d_v) │
│ │
│ 以 Llama-3-8B 为例(GQA: 32 KV heads, d_k=128): │
│ L = 32 层, n_kv_heads = 8, d_head = 128 │
│ │
│ 每层的 KV Cache: │
│ K: 1 × 8 × seq_len × 128 × 2 bytes (FP16) │
│ V: 1 × 8 × seq_len × 128 × 2 bytes │
│ = 2 × 8 × seq_len × 128 × 2 │
│ = 4096 × seq_len bytes │
│ │
│ 32 层的总 KV Cache: │
│ 32 × 4096 × seq_len = 131,072 × seq_len bytes │
│ │
│ 当 seq_len = 4096: │
│ 131,072 × 4096 = 537 MB │
│ │
│ 当 seq_len = 32768(长文本): │
│ 131,072 × 32768 = 4.3 GB │
│ │
│ 当 seq_len = 131072(超长文本): │
│ 131,072 × 131072 = 17.2 GB │
│ │
│ KV Cache 随序列长度线性增长,是长文本推理的主要瓶颈。 │
│ │
└─────────────────────────────────────────────────────────────┘2.5 采样策略——如何从概率分布中选下一个 token
模型输出 logits 后,经过 Softmax 得到一个概率分布。如何从分布中选择下一个 token,直接决定了生成文本的质量和多样性。
┌─────────────────────────────────────────────────────────────┐
│ 四种常见的采样策略 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 假设模型输出以下概率分布(简化,只显示前 5 个): │
│ │
│ token: "cat" "dog" "the" "a" "run" ... │
│ prob: 0.35 0.25 0.15 0.10 0.05 ... │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 1. Greedy(贪心) │ │
│ │ 永远选概率最大的 token:"cat" │ │
│ │ 优点:确定性强,速度快 │ │
│ │ 缺点:容易重复,缺乏多样性 │ │
│ │ 使用场景:代码生成、翻译(需要精确输出) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 2. Temperature Sampling(温度采样) │ │
│ │ logits_new = logits / T │ │
│ │ T → 0:趋近于 Greedy(极端分布) │ │
│ │ T = 0.7:常用默认值 │ │
│ │ T → ∞:趋近于均匀分布(完全随机) │ │
│ │ │ │
│ │ T=0.5: prob → [0.50, 0.30, 0.10, 0.05, 0.02] │ │
│ │ T=2.0: prob → [0.20, 0.18, 0.17, 0.16, 0.14] │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 3. Top-K Sampling │ │
│ │ 只从概率最高的 K 个 token 中采样 │ │
│ │ K=3: 只在 {"cat":0.35, "dog":0.25, "the":0.15} │ │
│ │ 中重新归一化后采样 │ │
│ │ 优点:避免选中极低概率的"垃圾"token │ │
│ │ 缺点:K 固定,不适应不同分布的形状 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 4. Top-P (Nucleus) Sampling │ │
│ │ 从累积概率 ≤ p 的最小 token 集合中采样 │ │
│ │ p=0.9: │ │
│ │ cat(0.35) + dog(0.25) + the(0.15) + a(0.10) │ │
│ │ + run(0.05) = 0.90 │ │
│ │ → 从这 5 个 token 中采样 │ │
│ │ 优点:动态调整候选集大小,适应不同分布 │ │
│ │ 缺点:计算稍复杂 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 现代 LLM 常用组合:Temperature + Top-P │
│ 比如 OpenAI 默认:T=1.0, Top-P=1.0(实际使用时有微调) │
│ │
└─────────────────────────────────────────────────────────────┘2.6 推理的完整伪代码
def generate(model, prompt_ids, max_new_tokens, temperature, top_p):
"""
自回归生成的核心循环。
"""
# 初始化:只有 prompt token
input_ids = prompt_ids.copy()
past_key_values = None # KV Cache,初始为空
for step in range(max_new_tokens):
# 前向传播(只算新 token 或全部 token)
if past_key_values is None:
# Prefill 阶段:并行处理所有 prompt token
logits, past_key_values = model.forward(
input_ids,
use_cache=True
)
else:
# Decode 阶段:只处理最后一个新 token
logits, past_key_values = model.forward(
input_ids[:, -1:], # 只取最后一个 token
past_key_values=past_key_values, # 复用 KV Cache
use_cache=True
)
# 取最后一个位置的 logits
next_logits = logits[:, -1, :] # (1, vocab_size)
# 应用温度
next_logits = next_logits / temperature
# 转为概率
probs = softmax(next_logits)
# Top-P 过滤
probs = top_p_filter(probs, top_p)
# 采样
next_token = sample_from(probs)
# 追加到序列
input_ids = concat([input_ids, next_token])
# 检查终止条件
if next_token == eos_token_id:
break
return input_ids第3部分:Compute Comparison——FLOPs 的量化对比
3.1 一次前向传播的 FLOPs 估算
Transformer 模型的计算量主要由两个部分组成:Attention 和 FFN。
┌─────────────────────────────────────────────────────────────┐
│ Transformer 单层 FLOPs 估算(简化公式) │
├─────────────────────────────────────────────────────────────┤
│ │
│ 符号定义: │
│ n = 序列长度(sequence length) │
│ d = 隐藏维度(hidden dim, d_model) │
│ f = FFN 中间维度(通常 = 4d 或 8d/3) │
│ V = 词表大小 │
│ │
│ 1. QKV 投影:3 × 2 × n × d × d = 6nd² │
│ (每个投影是 n×d × d×d,乘以 2 是乘法和加法各算一次) │
│ │
│ 2. Attention Score:2 × n × d × n = 2n²d │
│ (Q × K^T: n×d × d×n = n²d,加乘法各一次) │
│ │
│ 3. Attention Output:2 × n × n × d = 2n²d │
│ (Attn × V: n×n × n×d = n²d) │
│ │
│ 4. Output 投影:2 × n × d × d = 2nd² │
│ │
│ 5. FFN 第一层:2 × n × d × f = 2ndf │
│ FFN 第二层:2 × n × f × d = 2ndf │
│ FFN 总计:4ndf │
│ │
│ 单层总计(Attention + FFN): │
│ ≈ 8nd² + 4n²d + 4ndf │
│ (当 f = 4d 时)≈ 8nd² + 4n²d + 16nd² = 24nd² + 4n²d │
│ │
│ 全部 L 层:L × (24nd² + 4n²d) │
│ │
│ LM Head(输出投影):2 × n × d × V │
│ │
└─────────────────────────────────────────────────────────────┘3.2 训练 vs 推理的计算量对比表
┌─────────────────────────────────────────────────────────────┐
│ Training vs Inference 计算量对比 │
├─────────────────────────────────────────────────────────────┤
│ │
│ │ │ Training │ Inference │
│ ├─────────────────────┼──────────────────┼─────────────────┤
│ │ 前向传播次数 │ 1 │ N(每生成1个token│
│ │ │ │ 做1次前向传播) │
│ │ 每次前向的序列长度 │ 固定(如2048) │ 从1增长到N │
│ │ 并行度 │ 全序列并行 │ 逐token串行 │
│ │ 反向传播 │ 是(约2倍前向) │ 否 │
│ │ 批量大小 │ 大(百万级tokens) │ 1(或小batch) │
│ │ 梯度累积/通信 │ 有 │ 无 │
│ │ │ │ │
│ │ 总FLOPs(估) │ ~6× 一次前向 │ ~2× 一次前向/步 │
│ │ 总FLOPs(N tokens) │ ~6× (24Lnd²+...) │ ~N× (24Ld²+...) │
│ │ │ │ │
│ │ 瓶颈 │ 计算+通信 │ 显存带宽 │
│ │ │ (compute-bound) │ (memory-bound) │
│ │
└─────────────────────────────────────────────────────────────┘3.3 为什么训练一次前向约等于 3 次推理前向
反向传播需要计算所有中间变量的梯度。对于矩阵乘法 Y = X × W:
前向:Y = X × W → 1 次矩阵乘法
反向:
dL/dX = dL/dY × W^T → 1 次矩阵乘法
dL/dW = X^T × dL/dY → 1 次矩阵乘法
总计:前向 1 + 反向 2 = 3 次等效矩阵乘法实际中因为 activation functions、layer norm 等的额外计算,反向传播的成本大约是前向传播的 2~3 倍。所以训练一次迭代的总计算量 ≈ 前向 + 反向 ≈ 1 + 2.5 ≈ 3.5 倍单次前向。
如果使用 activation checkpointing(只保存部分激活值,需要时重新计算),还要额外增加一次"重新计算的前向传播"——但这可以显著节省显存。
┌─────────────────────────────────────────────────────────────┐
│ Activation Checkpointing 的权衡 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 不保存中间激活 → 显存占用 ↓ 但计算量 ↑ │
│ │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ Full Activations: │ │
│ │ 计算:Forward(1×) + Backward(2×) = 3× │ │
│ │ 显存:保存 ALL 激活值 │ │
│ │ │ │
│ │ Checkpoint EVERY layer: │ │
│ │ 计算:Forward(1×) + Recompute(1×) + Backward(2×) │ │
│ │ = 4× │ │
│ │ 显存:只保存每层输出 │ │
│ │ │ │
│ │ 额外计算 +33%,显存节省 ~70% │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 几乎所有大模型训练都使用 activation checkpointing。 │
│ 因为显存比计算更稀缺。 │
│ │
└─────────────────────────────────────────────────────────────┘3.4 推理的两个阶段:Prefill vs Decode
现代 LLM 推理服务(如 vLLM、TGI)通常将推理分为两个阶段:
┌─────────────────────────────────────────────────────────────┐
│ Prefill 阶段 vs Decode 阶段 │
├─────────────────────────────────────────────────────────────┤
│ │
│ Prefill(预填充/编码): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 输入:完整的 prompt(如用户问题,可能 500 tokens) │ │
│ │ 行为:一次性处理所有 prompt tokens(并行) │ │
│ │ 产出: │ │
│ │ - 最后一个位置的 logits(用于生成第一个新 token) │ │
│ │ - 所有位置的 KV Cache │ │
│ │ 计算特点:compute-bound(大量并行矩阵乘法) │ │
│ │ 和训练的前向传播几乎一样。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ Decode(解码/生成): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 输入:每次一个新 token │ │
│ │ 行为:串行生成,每次一个 token │ │
│ │ 使用 KV Cache 避免重复计算 │ │
│ │ 计算特点:memory-bound │ │
│ │ - 每次只做一个小矩阵乘法 │ │
│ │ - 大部分时间花在从显存读取模型权重和 KV Cache │ │
│ │ - GPU 计算单元大量闲置 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 这是推理优化的核心洞察: │
│ Prefill 阶段 → GPU 计算单元饱和(compute-bound) │
│ Decode 阶段 → GPU 计算单元空闲,显存带宽是瓶颈(mem-bound)│
│ │
│ 所以推理优化的重点是提高显存带宽利用率,不是堆算力。 │
│ │
└─────────────────────────────────────────────────────────────┘第4部分:Memory Comparison——为什么训练吃显存如喝水
4.1 显存占用的四大来源
┌─────────────────────────────────────────────────────────────┐
│ 模型运行时的显存占用四部分 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 1. 模型权重 (Model Weights) │
│ 训练和推理都需要。 │
│ FP32: 参数量 × 4 bytes │
│ FP16: 参数量 × 2 bytes │
│ INT8: 参数量 × 1 byte │
│ INT4: 参数量 × 0.5 bytes │
│ │
│ 2. 优化器状态 (Optimizer States) —— 仅训练 │
│ AdamW 为每个参数存储两个状态: │
│ m (一阶动量): 参数量 × 4 bytes (FP32) │
│ v (二阶动量): 参数量 × 4 bytes (FP32) │
│ 合计: 参数量 × 8 bytes │
│ │
│ 3. 梯度 (Gradients) —— 仅训练 │
│ 每个参数一个梯度值: │
│ FP32: 参数量 × 4 bytes │
│ │
│ 4. 激活值 (Activations) —— 训练 > 推理 │
│ 训练:需要保存所有中间激活用于反向传播 │
│ 推理:只需要当前层的激活,用完即丢 │
│ 但推理有 KV Cache(见 2.4 节)。 │
│ │
└─────────────────────────────────────────────────────────────┘4.2 具体例子:7B 模型的训练 vs 推理显存
┌─────────────────────────────────────────────────────────────┐
│ Llama-3-8B (实际 ~7B 参数) 的显存分析 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 模型配置: │
│ d_model = 4096, L = 32, vocab_size = 128256 │
│ 参数总量 ≈ 8.03B(含 embedding 和 LM head) │
│ │
│ ═══════════ 训练显存 ═══════════ │
│ │
│ 模型权重 (FP32 或 mixed-precision): │
│ 8B × 4 bytes = 32 GB │
│ │
│ 优化器状态 (AdamW, FP32): │
│ 8B × 8 bytes = 64 GB ← 比模型本身还大! │
│ │
│ 梯度 (FP32): │
│ 8B × 4 bytes = 32 GB │
│ │
│ 激活值(取决于 batch_size 和 seq_len): │
│ 假设 batch=1, seq_len=4096, 使用 activation ckpt: │
│ ≈ 15-25 GB │
│ │
│ 训练总计:32 + 64 + 32 + 20 = ~148 GB │
│ │
│ → 一块 H100 (80GB) 装不下,需要至少 2 块。 │
│ → 实际上常用 8×H100 做分布式训练。 │
│ │
│ ═══════════ 推理显存 ═══════════ │
│ │
│ 模型权重 (FP16): │
│ 8B × 2 bytes = 16 GB │
│ │
│ 优化器状态:无 (0 GB) │
│ │
│ 梯度:无 (0 GB) │
│ │
│ 激活值(用完即丢): │
│ ≈ 0.5-1 GB │
│ │
│ KV Cache (seq_len=4096, FP16): │
│ ≈ 0.5 GB(见 2.4 节计算) │
│ │
│ 推理总计:16 + 0 + 0 + 1 + 0.5 = ~17.5 GB │
│ │
│ → 一块 24GB 显卡(如 RTX 4090)就能跑。 │
│ → 如果用 INT4 量化:8B × 0.5 = 4GB,总计 ~6GB。 │
│ │
│ ═══════════ 对比 ═══════════ │
│ │
│ │ │ 训练 (FP32) │ 推理 (FP16) │ │
│ ├──────────────┼────────────────┼────────────────┤ │
│ │ 模型权重 │ 32 GB │ 16 GB │ │
│ │ 优化器状态 │ 64 GB │ 0 GB │ │
│ │ 梯度 │ 32 GB │ 0 GB │ │
│ │ 激活值/KV │ ~20 GB │ ~1.5 GB │ │
│ │ 总计 │ ~148 GB │ ~17.5 GB │ │
│ │ 比例 │ ~8.5× │ 1× │ │
│ │
│ 训练需要的显存大约是推理的 8-10 倍。 │
│ 其中最大的"浪费"是优化器状态(训练特有)。 │
│ │
└─────────────────────────────────────────────────────────────┘4.3 为什么量化对推理特别有效
┌─────────────────────────────────────────────────────────────┐
│ 量化对训练和推理的不同意义 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 量化 = 用低精度表示权重(如 FP16 → INT8 → INT4) │
│ │
│ 对推理的影响(巨大): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 推理是 memory-bound。 │ │
│ │ 减小模型 = 更少的数据从显存搬运到计算单元。 │ │
│ │ │ │
│ │ FP16 → INT4: 模型大小 ÷ 4 │ │
│ │ 7B 模型:16GB → 4GB │ │
│ │ 原来需要 24GB 显卡,现在 8GB 就能跑。 │ │
│ │ │ │
│ │ 而且推理不需要反向传播,精度损失影响较小。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 对训练的影响(有限): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 训练是 compute-bound 且需要精确梯度。 │ │
│ │ 权重可以用 FP16 存储(mixed-precision training), │ │
│ │ 但优化器状态和梯度仍然需要 FP32。 │ │
│ │ │ │
│ │ 训练时通常不能直接使用 INT8/INT4 权重—— │ │
│ │ 梯度更新需要高精度,量化误差会累积。 │ │
│ │ │ │
│ │ QLoRA 等方法是曲线救国: │ │
│ │ 主模型用 INT4(冻结),只训练少量 LoRA 参数(FP32)。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 这也是为什么推理可以用一台 MacBook 跑 7B 模型, │
│ 但训练 7B 模型需要几台 H100 的原因之一。 │
│ │
└─────────────────────────────────────────────────────────────┘第5部分:为什么训练和推理的优化方向不同
5.1 两个完全不同的优化目标
┌─────────────────────────────────────────────────────────────┐
│ 训练优化 vs 推理优化 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 训练优化目标: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 核心指标:Throughput(吞吐量) │ │
│ │ tokens/second across all GPUs │ │
│ │ │ │
│ │ 不关心的指标:单个样本的 Latency(延迟) │ │
│ │ → 训练一个 epoch 需要几小时甚至几天, │ │
│ │ 单个 step 是 0.1 秒还是 0.5 秒差别不大。 │ │
│ │ │ │
│ │ 优化手段: │ │
│ │ • 大 batch size(提高 GPU 利用率) │ │
│ │ • 数据并行 / 模型并行 / 流水线并行 │ │
│ │ • 梯度累积(模拟大 batch) │ │
│ │ • Mixed-precision training(FP16 + FP32) │ │
│ │ • Flash Attention(节省显存,允许更大 batch) │ │
│ │ • Activation Checkpointing(同上) │ │
│ │ • ZeRO(分布式优化器状态) │ │
│ │ • 高速互联(NVLink, InfiniBand) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 推理优化目标: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 核心指标1:TTFT(Time To First Token,首token延迟) │ │
│ │ → 用户发问到看到第一个字的时间 │ │
│ │ → 对应 Prefill 阶段 │ │
│ │ │ │
│ │ 核心指标2:TPOT(Time Per Output Token,每token延迟)│ │
│ │ → 生成过程中每个 token 的时间 │ │
│ │ → 对应 Decode 阶段 │ │
│ │ │ │
│ │ 核心指标3:Throughput(吞吐量,tokens/second) │ │
│ │ → 同时服务多个用户时的总体生成速度 │ │
│ │ │ │
│ │ 优化手段: │ │
│ │ • KV Cache(避免重复计算) │ │
│ │ • 量化(INT8/INT4,减小模型,提高带宽利用率) │ │
│ │ • Flash Attention(减少 KV Cache 显存) │ │
│ │ • Speculative Decoding(用草稿模型"猜测"多个token)│ │
│ │ • Continuous Batching(动态拼接请求) │ │
│ │ • PagedAttention / vLLM(KV Cache 分页管理) │ │
│ │ • 高显存带宽显卡(HBM3 > GDDR6X > GDDR6) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────┘5.2 为什么训练和推理使用不同的硬件
┌─────────────────────────────────────────────────────────────┐
│ 训练硬件 vs 推理硬件 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 训练硬件(如 NVIDIA H100): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 重点:算力 + 互联 │ │
│ │ │ │
│ │ • 大量 Tensor Core(高 FP16/BF16/FP8 TFLOPS) │ │
│ │ • NVLink/NVSwitch:GPU 间高速互联(900 GB/s) │ │
│ │ • HBM3 高带宽显存(3.35 TB/s) │ │
│ │ • 大显存(80GB HBM3) │ │
│ │ │ │
│ │ 为什么需要高速互联? │ │
│ │ → 分布式训练中 GPU 间频繁通信梯度 │ │
│ │ → 慢互联 = GPU 空等 = 浪费算力 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 推理硬件(如 NVIDIA L40S / T4 / A10): │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 重点:显存带宽 + 成本 │ │
│ │ │ │
│ │ • 适中的算力(推理 Decode 阶段算力需求低) │ │
│ │ • 高显存带宽(读取模型权重和 KV Cache 的瓶颈) │ │
│ │ • 足够大的显存(装下模型 + KV Cache) │ │
│ │ • 低功耗(数据中心电费是长期成本) │ │
│ │ • 不需要 NVLink(推理不需要 GPU 间通信) │ │
│ │ │ │
│ │ L40S vs H100 推理对比: │ │
│ │ L40S: 48GB GDDR6, 带宽 864 GB/s, 功耗 350W │ │
│ │ H100: 80GB HBM3, 带宽 3350 GB/s, 功耗 700W │ │
│ │ │ │
│ │ 对推理来说 L40S 性价比远高于 H100: │ │
│ │ → Decode 阶段根本用不满 H100 的算力 │ │
│ │ → H100 的高速互联在推理中完全浪费 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 苹果 M 系列芯片的推理优势: │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ • 统一内存架构:CPU 和 GPU 共享大容量内存 │ │
│ │ M2 Ultra: 192GB 统一内存 │ │
│ │ • 可以跑 FP16 的 70B 模型 │ │
│ │ • 内存带宽(800 GB/s)虽不如 HBM3 但足够推理 │ │
│ │ • 不适合训练(算力不足,没有 CUDA 生态) │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────┘5.3 推理优化的前沿技术
┌─────────────────────────────────────────────────────────────┐
│ 推理优化的前沿技术概览 │
├─────────────────────────────────────────────────────────────┤
│ │
│ 1. Speculative Decoding(投机解码) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 问题:大模型一次只能生成一个 token,慢。 │ │
│ │ │ │
│ │ 方案:用一个小"草稿模型"(draft model) 快速生成 │ │
│ │ K 个候选 token,然后大模型一次性验证这 K 个。 │ │
│ │ │ │
│ │ 草稿模型(快但不准):guess → "I" "love" "you" │ │
│ │ 大模型(慢但准):并行验证这三个 → 全对! │ │
│ │ │ │
│ │ 一次大模型前向传播 → 接受/拒绝 K 个 token。 │ │
│ │ 理想情况:吞吐量 × K。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 2. Continuous Batching(连续批处理) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 问题:不同用户请求长度不一,有些早结束有些还在生成。 │ │
│ │ │ │
│ │ 传统方案:等整个 batch 完成再一起返回。 │ │
│ │ → 早完成的请求在等,GPU 利用率低。 │ │
│ │ │ │
│ │ Continuous Batching: │ │
│ │ 请求完成 → 立刻从 batch 中移除 → 加入新请求。 │ │
│ │ 不像传统 batch 那样"等齐了再走"。 │ │
│ │ │ │
│ │ 效果:GPU 利用率从 ~30% 提升到 ~80%。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 3. PagedAttention / vLLM │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 问题:KV Cache 的显存管理低效。 │ │
│ │ → 预分配固定大小的 KV Cache,大量浪费。 │ │
│ │ → 类似操作系统的"内部碎片"问题。 │ │
│ │ │ │
│ │ 方案:借鉴操作系统虚拟内存的分页机制。 │ │
│ │ KV Cache 分成固定大小的"页"(blocks)。 │ │
│ │ 按需分配和回收,不连续存储也没关系。 │ │
│ │ │ │
│ │ 效果:KV Cache 显存利用率从 ~30% 提升到 ~96%。 │ │
│ │ → 同样的显存服务更多请求。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
│ 4. Flash Attention(对训练和推理都有用) │
│ ┌─────────────────────────────────────────────────────┐ │
│ │ 问题:标准 Attention 需要 O(n²) 显存存 Score 矩阵。 │ │
│ │ │ │
│ │ 方案:分块计算(tiling)+ 在线 Softmax。 │ │
│ │ 不把完整的 n×n Score 矩阵写回 HBM(显存)。 │ │
│ │ 在 SRAM(片上缓存)里分块算完。 │ │
│ │ │ │
│ │ 效果: │ │
│ │ 训练:节省激活值显存 → 允许更大 batch。 │ │
│ │ 推理:减少 KV Cache 的中间数据搬运。 │ │
│ └─────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────┘核心总结
总结1:训练和推理的本质差异
| 维度 | 训练 (Training) | 推理 (Inference) |
|---|---|---|
| 输入 | 完整目标序列(Teacher Forcing) | 逐 token 生成(Auto-Regressive) |
| 并行度 | 全序列并行 | 逐 token 串行 |
| 反向传播 | 有(~2倍前向计算量) | 无 |
| 权重 | 不断更新 | 冻结不变 |
| 目标 | 最小化 Loss | 生成合理文本 |
| 计算瓶颈 | 算力(compute-bound) | 显存带宽(memory-bound) |
| 显存瓶颈 | 激活值 + 优化器状态 | KV Cache + 模型权重 |
总结2:Teacher Forcing 让并行训练成为可能
- 输入完整序列 + Causal Mask = 所有位置并行计算但互不"偷看"
- 每个位置独立计算 Loss,总 Loss 是所有位置的平均
- 没有 Teacher Forcing,训练将和推理一样慢(N 次串行前向传播)
总结3:KV Cache 是推理优化的灵魂
- 核心观察:已有 token 的 K 和 V 不会因新 token 而改变
- 缓存 K 和 V,每步只算新 token → 计算量从 O(N³) 降到 O(N²)
- 代价:额外显存(随序列长度线性增长)
- 长文本推理的瓶颈就是 KV Cache 的显存
总结4:训练显存 >> 推理显存
- 7B 模型训练:~148 GB,主要是优化器状态(64GB)和梯度(32GB)
- 7B 模型推理:~17.5 GB(FP16),量化后 ~6 GB(INT4)
- 训练显存约是推理的 8-10 倍
- 量化对推理特别有效(memory-bound),对训练作用有限(需要高精度梯度)
总结5:训练和推理的优化方向完全不同
- 训练:拼吞吐(tokens/s),需要强算力 + 高速 GPU 互联。硬件 = H100 集群。
- 推理:拼延迟(TTFT + TPOT)和吞吐,需要高显存带宽。硬件 = L40S / T4。
- Prefill 像训练前向(compute-bound),Decode 是推理独有(memory-bound)
- 推理优化(量化、投机解码、Continuous Batching)追求的是"用更少的显存服务更多用户"
章节测试
测试1:Teacher Forcing 中,Decoder 的输入是什么?
A. 模型自己上一步预测的 token B. 完整的目标序列(正确答案),通过 Causal Mask 防止看到未来 C. 只有起始标记 <s> D. 随机采样的 token 序列
测试2:训练时所有位置可以并行计算,但不会"偷看"未来 token。这是通过什么机制实现的?
A. 把未来 token 的 embedding 设为 0 B. Causal Mask(Attention 分数矩阵的上三角设为 -∞) C. 每个位置独立做一次前向传播 D. 反向传播时修正
测试3:推理时使用 KV Cache 后,生成第 N 个 token 时的 Attention 计算复杂度是多少?
A. O(1) —— 常数时间 B. O(N) —— 和当前序列长度成线性关系 C. O(N²) —— 和当前序列长度的平方成线性关系 D. O(N³) —— 和当前序列长度的立方成线性关系
测试4:7B 模型训练需要约 150GB 显存,其中最大的单一来源是什么?
A. 模型权重(~32GB) B. 优化器状态——Adam 的 m 和 v(~64GB) C. 梯度(~32GB) D. 激活值(~20GB)
测试5:为什么推理的 Decode 阶段是 memory-bound 而不是 compute-bound?
测试6:为什么量化(如 INT4)对推理的加速效果远大于对训练的加速效果?
测试7:简述 Teacher Forcing 和 Auto-Regressive 的区别,并解释为什么推理不能使用 Teacher Forcing。
参考答案
测试1答案
答案:B。Teacher Forcing 的核心就是"用正确答案作为输入"。Decoder 接收到完整的目标序列,通过 Causal Mask 确保每个位置只能看到当前及之前的 token。这样所有位置可以在一次前向传播中并行计算。
测试2答案
答案:B。Causal Mask 是一个上三角为 -∞ 的矩阵,加在 Attention Scores 上。Softmax 后 -∞ 位置的权重为 0,因此每个位置只能 attend 到自己及之前的 token。这是"并行计算但因果保序"的数学基础。
测试3答案
答案:B。使用 KV Cache 后,第 N 步只需要计算新 token 的 Q(1×d)与缓存的 K_cache(N×d)的乘积,复杂度为 O(N)(准确的说是 O(N×d))。而没有 KV Cache 时是 O(N²)。注意:总 N 步的累积复杂度仍是 O(N²),但比无缓存的 O(N³) 好得多。
测试4答案
答案:B。AdamW 优化器为每个参数维护两个状态 m(一阶动量)和 v(二阶动量),都是 FP32(4 bytes)。总大小 = 2 × 参数量 × 4 = 参数量 × 8 bytes。对 8B 参数模型就是 64 GB,比模型本身(32 GB FP32)还大一倍。
测试5答案
推理的 Decode 阶段每次只处理一个 token。由于序列长度短(KV Cache 中的历史 token 的 K 和 V 只是被读取),矩阵乘法非常小。GPU 的 Tensor Core 大部分时间在等待数据从 HBM(显存)搬运到片上 SRAM。所以瓶颈不是计算速度,而是显存带宽——这就是 memory-bound。Prefill 阶段处理整个 prompt(可能数千 token),矩阵乘法大,GPU 计算单元饱和,是 compute-bound。
测试6答案
推理是 memory-bound——瓶颈在显存带宽而不在算力。量化直接缩小模型(INT4 = 1/8 FP32),同样的显存带宽下每秒能搬运更多参数 → 推理加速。训练是 compute-bound——需要大量矩阵乘法算力,而且反向传播需要高精度梯度(量化误差会累积导致训练不稳定)。虽然可以用 mixed-precision(FP16 前向 + FP32 优化器),但不能像推理那样激进量化。
测试7答案
Teacher Forcing(训练):输入完整的目标序列(正确答案),所有位置并行计算,一次前向传播得到所有位置的预测。 Auto-Regressive(推理):每次只生成一个 token,用自己的预测作为下一步输入,N 个 token 需要 N 次前向传播。 推理不能使用 Teacher Forcing 的原因:推理时没有"正确答案"——你不知道目标序列是什么。你只有一个起始标记,必须先生成第一个 token 才知道第二个位置的输入是什么。Teacher Forcing 要求提前知道整个序列,这在生成任务中是不可能的。
相关笔记
- [[08-transformer-by-hand]] — 手动计算 Transformer 的完整数据流(本文依赖的前置知识)
- [[07-llm-evolution]] — 从 Word2Vec 到 Transformer 的历史演进
- [[09-decoder-only-llm]] — GPT 为什么只需要 Decoder(不需要 Encoder)
- [[11-distributed-training]] — 分布式训练:数据并行、模型并行、ZeRO
- [[12-inference-optimization]] — 推理优化深入:Flash Attention、vLLM、量化
下一步学习
- [ ] 用 PyTorch 实现一个带 KV Cache 的自回归生成循环,对比有/无 Cache 的速度差异
- [ ] 阅读 vLLM 论文 (PagedAttention) —— 理解 KV Cache 的分页管理
- [ ] 阅读 Flash Attention 论文 —— 理解 IO-aware Attention 计算
- [ ] 实际测一下:用 llama.cpp 加载同一个模型,对比 FP16 / INT8 / INT4 的推理速度和显存占用
- [ ] 思考:如果你有 4 张 24GB 显卡,你如何训练一个 7B 模型?(提示:ZeRO、模型并行、梯度累积)
学习状态:🟡 开始学习