Skip to content
Gains Summary
Main Navigation 首页 / Home
C++ 编程 / C++ Programming
系统与高性能 / Systems & Performance
Web 开发 / Web Development
人工智能 / Artificial Intelligence
工业软件 / Industrial Software
其他内容 / Other Topics
C++ 编程 / C++系统与性能 / SystemsWeb 开发 / Web人工智能 / AI工业软件 / Industrial

外观

Sidebar Navigation

← 人工智能 / Artificial Intelligence

大语言模型 / Large Language Models

1. LLM 前置知识学习路线 / A Prerequisite Learning Path for Large Language Models

2. 神经网络基础 - 从零理解 AI 的"计算单元" / Neural Network Fundamentals from Artificial Neurons

3. Token 与上下文窗口 - LLM 的计费与记忆单位 / Tokens and Context Windows as the Units of LLM Cost and Memory

4. 解码策略 - 控制 LLM 输出的艺术 / Decoding Strategies for Controlling LLM Output

5. 消息角色 - 构建 Agent 对话的基础 / Message Roles as the Foundation of Agent Conversations

6. 流式输出 - 实时交互的体验优化 / Streaming Output for Responsive Interaction

7. Prompt 工程基础 - 与 LLM 高效对话的技巧 / Prompt Engineering Fundamentals for Effective LLM Interaction

8. LLM 进化史 - 从词向量到 Transformer / The Evolution of LLMs from Word Embeddings to Transformers

9. Transformer 手动计算:从 Attention 到 Encoder-Decoder 完整数据流

10. Transformer 核心原理 - 现代 LLM 的基石 / Transformer Fundamentals Behind Modern LLMs

11. Decoder-Only LLM 深度解析:为什么扔掉 Encoder,以及 KV Cache 如何工作

12. 训练 vs 推理:同一个 Transformer,两条完全不同的执行路径

13. Transformer 训练阶段计算详解 - 手算每一行矩阵 / Transformer Training Computation Matrix by Matrix

14. Transformer 推理阶段详解 — 模型如何"思考"并生成回答 / Transformer Inference and Autoregressive Generation

15. 训练基础扫盲 - 理解 Fine-tune 在做什么 / A Training Primer for Understanding Fine-Tuning

16. LLM 预训练全景:数据管道、Scaling Laws 与训练稳定性

17. 训练基础设施 - 从单卡到千卡集群 / Training Infrastructure from One GPU to Thousand-GPU Clusters

18. Post-Training Pipeline - 从 Base Model 到可用助手 / The Post-Training Pipeline from Base Model to Assistant

19. SFT 深度解析:从 Base Model 到指令跟随——后训练第一步 / SFT Deep Dive: Teaching Base Models to Follow Instructions

20. RLHF 深度解析:从 Reward Model 到 PPO 的完整对齐流程

21. DPO 与对齐方法:从 RLHF 复杂度到直接偏好优化

22. 研究视角:DL/RL 理论到 LLM 训练的完整映射

本页目录

训练 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 让训练"假装"并行                  │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43

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
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28

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]
1
2
3

准备训练数据——右移一位:

┌─────────────────────────────────────────────────────────────┐
│              训练数据的构造: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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25

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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50

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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18

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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30

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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59

1.7 为什么训练时激活值必须保存 ​

反向传播需要计算 dL/dW。以最简单的线性层 Y = X × W 为例:

前向:Y = X × W
反向:dL/dW = X^T × (dL/dY)

要计算 dL/dW,你需要:
- dL/dY(从上一层反向传回来,已知)
- X(前向传播时的输入,必须保存!)

如果你没保存 X,就得重新算一次前向传播——这就是 gradient checkpointing 的思想。
1
2
3
4
5
6
7
8

对于 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 次前向传播。                         │
│  这就是为什么"训练一个模型"很快,但"用模型生成文本"很慢。    │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38

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。
  ...
1
2
3
4
5
6
7
8
9
10
11
12

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 倍。                                         │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37

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 随序列长度线性增长,是长文本推理的主要瓶颈。        │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32

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(实际使用时有微调)     │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52

2.6 推理的完整伪代码 ​

python
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
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47

第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                         │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34

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)  │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21

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 次等效矩阵乘法
1
2
3
4
5
6

实际中因为 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。         │
│  因为显存比计算更稀缺。                                      │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23

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)│
│                                                             │
│  所以推理优化的重点是提高显存带宽利用率,不是堆算力。         │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33

第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 节)。                         │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27

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 倍。                        │
│  其中最大的"浪费"是优化器状态(训练特有)。                  │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63

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 的原因之一。                    │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35

第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)         │   │
│  └─────────────────────────────────────────────────────┘   │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48

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 生态)             │   │
│  └─────────────────────────────────────────────────────┘   │
│                                                             │
└─────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47

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
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60

核心总结 ​

总结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、模型并行、梯度累积)

学习状态:🟡 开始学习

最后更新于:

Pager
上一篇11. Decoder-Only LLM 深度解析:为什么扔掉 Encoder,以及 KV Cache 如何工作
下一篇13. Transformer 训练阶段计算详解 - 手算每一行矩阵 / Transformer Training Computation Matrix by Matrix

持续记录,持续成长

Copyright © Tidenflow