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 训练的完整映射

本页目录

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

📅 创建时间:2026-04-29 🏷️ 标签:#Transformer #训练 #QKV #Multi-Head #Temperature #采样策略 📚 前置知识:[[01 - 神经网络基础]] [[07 - LLM 进化史]] 🎯 文档类型:计算详解 · 手把手数值推导


📋 文档目标 ​

阅读完本文档后,你将能够:

  • [ ] 手算 Embedding + 位置编码的完整过程
  • [ ] 独立完成一次 Self-Attention 的前向传播(Q/K/V → Score → Softmax → Output)
  • [ ] 理解单头注意力的局限性,以及多头如何解决
  • [ ] 完整追踪 Multi-Head Attention 的矩阵拼接与投影
  • [ ] 理解从 Logits → 概率分布 → Temperature 调参的全链路
  • [ ] 区分 TopK 和 TopP 采样,并理解为何通常不同时使用
  • [ ] 建立从 token 输入到 loss 计算的完整数据流直觉

第1章:从 Token 到向量 — 输入Embedding ​

1.1 我们的示例任务 ​

我们用最简单的句子来演示:

输入序列:"猫 追 老鼠"(3个Token)
1

为了手算方便,我们约定以下超参数:

d_model = 4      # 词向量维度(最终输出向量维度)
seq_len  = 3     # 序列长度(3个词)
d_vocab  = 6     # 假设词表大小为6
1
2
3

1.2 Token 到索引 ​

首先,每个词被映射为一个整数索引(词表中的位置):

词表(按字母排序,仅作示例):
索引 0: "<PAD>"  (填充符)
索引 1: "老鼠"
索引 2: "猫"
索引 3: "追"
索引 4: "我"
索引 5: "爱"

输入序列 "猫 追 老鼠" 对应的索引:
Token   索引
--------------
猫  →    2
追  →    3
老鼠 →    1
1
2
3
4
5
6
7
8
9
10
11
12
13
14

1.3 词嵌入(Embedding Lookup) ​

词嵌入层是一个形状为 (d_vocab, d_model) 的查找表。对于我们的例子:

嵌入矩阵 W_embed (6 × 4):

         dim 0   dim 1   dim 2   dim 3
索引 0: [ 0.00,  0.00,  0.00,  0.00]   # <PAD>
索引 1: [ 0.60,  0.10,  0.30,  0.90]   # 老鼠
索引 2: [ 0.50,  0.20,  0.80,  0.10]   # 猫
索引 3: [ 0.30,  0.70,  0.20,  0.60]   # 追
索引 4: [ 0.70,  0.50,  0.10,  0.40]   # 我
索引 5: [ 0.20,  0.90,  0.70,  0.30]   # 爱
1
2
3
4
5
6
7
8
9

通过索引查找,得到每个词的初始向量:

Token "猫"(索引2)→ X_猫 = [0.50, 0.20, 0.80, 0.10]
Token "追"(索引3)→ X_追 = [0.30, 0.70, 0.20, 0.60]
Token "老鼠"(索引1)→ X_老鼠 = [0.60, 0.10, 0.30, 0.90]

构成输入矩阵 X (3 × 4):
       dim0  dim1  dim2  dim3
位置0: [0.50, 0.20, 0.80, 0.10]   # 猫
位置1: [0.30, 0.70, 0.20, 0.60]   # 追
位置2: [0.60, 0.10, 0.30, 0.90]   # 老鼠
1
2
3
4
5
6
7
8
9

1.4 位置编码(Positional Encoding) ​

Attention 本身无法感知位置,因此需要显式注入位置信息。

我们使用原始 Transformer 论文中的 Sinusoidal 编码:

PE(pos, 2i)   = sin(pos / 10000^(2i/d_model))
PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))

其中:
- pos = 位置(0, 1, 2, ...)
- i   = 维度索引(0, 1)  [因为 d_model/2 = 2]
1
2
3
4
5
6

具体计算(d_model = 4,故 i = 0, 1):

位置 0 (pos = 0):
PE(0, 0) = sin(0 / 10000^0)     = sin(0)     = 0.0000
PE(0, 1) = cos(0 / 10000^0)     = cos(0)     = 1.0000
PE(0, 2) = sin(0 / 10000^2)     = sin(0)     = 0.0000
PE(0, 3) = cos(0 / 10000^2)     = cos(0)     = 1.0000
→ PE(0) = [0.0000, 1.0000, 0.0000, 1.0000]

位置 1 (pos = 1):
PE(1, 0) = sin(1 / 10000^0)     = sin(1)     ≈ 0.8415
PE(1, 1) = cos(1 / 10000^0)     = cos(1)     ≈ 0.5403
PE(1, 2) = sin(1 / 10000^2)     = sin(0.0001)≈ 0.0001
PE(1, 3) = cos(1 / 10000^2)     = cos(0.0001)≈ 0.9999
→ PE(1) = [0.8415, 0.5403, 0.0001, 0.9999]

位置 2 (pos = 2):
PE(2, 0) = sin(2 / 10000^0)     = sin(2)     ≈ 0.9093
PE(2, 1) = cos(2 / 10000^0)     = cos(2)     ≈ -0.4161
PE(2, 2) = sin(2 / 10000^2)     = sin(0.0002)≈ 0.0002
PE(2, 3) = cos(2 / 10000^2)     = cos(0.0002)≈ 0.9999
→ PE(2) = [0.9093, -0.4161, 0.0002, 0.9999]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

位置编码矩阵 PE (3 × 4):

       dim 0   dim 1   dim 2   dim 3
位置0: [0.0000, 1.0000, 0.0000, 1.0000]
位置1: [0.8415, 0.5403, 0.0001, 0.9999]
位置2: [0.9093,-0.4161, 0.0002, 0.9999]
1
2
3
4

1.5 最终输入 = 词嵌入 + 位置编码 ​

X_input = X + PE(逐元素相加)

X_input = X + PE:
       dim0       dim1       dim2       dim3
位置0: [0.50+0.00, 0.20+1.00, 0.80+0.00, 0.10+1.00] = [0.50, 1.20, 0.80, 1.10]
位置1: [0.30+0.84, 0.70+0.54, 0.20+0.00, 0.60+1.00] = [1.14, 1.24, 0.20, 1.60]
位置2: [0.60+0.91, 0.10-0.42, 0.30+0.00, 0.90+1.00] = [1.51,-0.32, 0.30, 1.90]

最终输入矩阵 X_input (3 × 4):
[[ 0.50,  1.20,  0.80,  1.10],
 [ 1.14,  1.24,  0.20,  1.60],
 [ 1.51, -0.32,  0.30,  1.90]]
1
2
3
4
5
6
7
8
9
10
11
12

这就是传入 Transformer 层的第一个输入。


第2章:QKV 矩阵 — 线性投影 ​

2.1 Q/K/V 的核心思想 ​

Query(查询): 我当前这个词想要查找什么信息?
Key(键)     : 序列中每个位置能提供什么信息?
Value(值)    : 每个位置的实际信息内容
1
2
3

这三个向量都是通过同一个输入向量 X_input 做不同的线性投影得到的:

Q = X_input · W_q     (Query 投影)
K = X_input · W_k     (Key 投影)
V = X_input · W_v     (Value 投影)
1
2
3

投影矩阵 W_q, W_k, W_v 是可学习的参数,形状都是 (d_model, d_k)。

2.2 定义权重矩阵 ​

为了手算方便,我们设定 d_k = d_v = 3(投影后的维度):

W_q (4 × 3):
       q0      q1      q2
    [[0.10,   0.20,   0.30],
     [0.40,   0.50,   0.60],
     [0.70,   0.80,   0.90],
     [0.15,   0.25,   0.35]]

W_k (4 × 3):
       k0      k1      k2
    [[0.05,   0.15,   0.25],
     [0.35,   0.45,   0.55],
     [0.65,   0.75,   0.85],
     [0.10,   0.20,   0.30]]

W_v (4 × 3):
       v0      v1      v2
    [[0.20,   0.30,   0.40],
     [0.50,   0.60,   0.70],
     [0.80,   0.90,   1.00],
     [0.10,   0.20,   0.30]]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

2.3 计算 Q 矩阵 ​

矩阵乘法:Q = X_input · W_q

X_input 的第0行 [0.50, 1.20, 0.80, 1.10] 与 W_q 的每一列做点积:

Q[0, 0] = 0.50×0.10 + 1.20×0.40 + 0.80×0.70 + 1.10×0.15
        = 0.05 + 0.48 + 0.56 + 0.165
        = 1.255

Q[0, 1] = 0.50×0.20 + 1.20×0.50 + 0.80×0.80 + 1.10×0.25
        = 0.10 + 0.60 + 0.64 + 0.275
        = 1.615

Q[0, 2] = 0.50×0.30 + 1.20×0.60 + 0.80×0.90 + 1.10×0.35
        = 0.15 + 0.72 + 0.72 + 0.385
        = 1.975
1
2
3
4
5
6
7
8
9
10
11

类似地计算位置1和位置2:

位置0 (猫):  Q_猫   = [1.255, 1.615, 1.975]
位置1 (追):  Q_追   = [1.088, 1.404, 1.732]  # 略去中间计算
位置2 (老鼠): Q_老鼠 = [0.967, 1.269, 1.582]  # 略去中间计算
1
2
3

Q 矩阵 (3 × 3):

       q0      q1      q2
位置0: [1.255, 1.615, 1.975]   # 猫的查询
位置1: [1.088, 1.404, 1.732]   # 追的查询
位置2: [0.967, 1.269, 1.582]   # 老鼠的查询
1
2
3
4

2.4 计算 K 矩阵 ​

矩阵乘法:K = X_input · W_k

K[0, 0] = 0.50×0.05 + 1.20×0.35 + 0.80×0.65 + 1.10×0.10
        = 0.025 + 0.42 + 0.52 + 0.11
        = 1.075

K[0, 1] = 0.50×0.15 + 1.20×0.45 + 0.80×0.75 + 1.10×0.20
        = 0.075 + 0.54 + 0.60 + 0.22
        = 1.435

K[0, 2] = 0.50×0.25 + 1.20×0.55 + 0.80×0.85 + 1.10×0.30
        = 0.125 + 0.66 + 0.68 + 0.33
        = 1.795
1
2
3
4
5
6
7
8
9
10
11

K 矩阵 (3 × 3):

       k0      k1      k2
位置0: [1.075, 1.435, 1.795]   # 猫的键
位置1: [0.926, 1.243, 1.569]   # 追的键
位置2: [0.826, 1.118, 1.419]   # 老鼠的键
1
2
3
4

2.5 计算 V 矩阵 ​

矩阵乘法:V = X_input · W_v

V[0, 0] = 0.50×0.20 + 1.20×0.50 + 0.80×0.80 + 1.10×0.10
        = 0.10 + 0.60 + 0.64 + 0.11
        = 1.45

V[0, 1] = 0.50×0.30 + 1.20×0.60 + 0.80×0.90 + 1.10×0.20
        = 0.15 + 0.72 + 0.72 + 0.22
        = 1.81

V[0, 2] = 0.50×0.40 + 1.20×0.70 + 0.80×1.00 + 1.10×0.30
        = 0.20 + 0.84 + 0.80 + 0.33
        = 2.17
1
2
3
4
5
6
7
8
9
10
11

V 矩阵 (3 × 3):

       v0      v1      v2
位置0: [1.450, 1.810, 2.170]   # 猫的值
位置1: [1.245, 1.555, 1.875]   # 追的值
位置2: [1.095, 1.380, 1.675]   # 老鼠的值
1
2
3
4

2.6 QKV 总结 ​

输入 X_input (3 × 4)
      ↓
      ├─→ X_input · W_q = Q (3 × 3)
      ├─→ X_input · W_k = K (3 × 3)
      └─→ X_input · W_v = V (3 × 3)

Q = [[1.255, 1.615, 1.975],
     [1.088, 1.404, 1.732],
     [0.967, 1.269, 1.582]]

K = [[1.075, 1.435, 1.795],
     [0.926, 1.243, 1.569],
     [0.826, 1.118, 1.419]]

V = [[1.450, 1.810, 2.170],
     [1.245, 1.555, 1.875],
     [0.967, 1.269, 1.582]]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17

第3章:注意力分数计算 — QK^T / √d_k ​

3.1 为什么是 Q · K^T? ​

Q 问"我需要什么信息",K 说"我能提供什么信息"。

两者的点积衡量了"当前词对序列中其他词的关注程度":

score(i, j) = Q[i] · K[j]^T = 第i个词的查询 与 第j个词的键 的相似度
1

计算 Q · K^T(3 × 3 矩阵):

Q · K^T = Q (3×3) × K^T (3×3)

(QK^T)[0,0] = Q_猫 · K_猫
            = 1.255×1.075 + 1.615×1.435 + 1.975×1.795
            = 1.349 + 2.318 + 3.545
            = 7.212

(QK^T)[0,1] = Q_猫 · K_追
            = 1.255×0.926 + 1.615×1.243 + 1.975×1.569
            = 1.162 + 2.007 + 3.099
            = 6.268

(QK^T)[0,2] = Q_猫 · K_老鼠
            = 1.255×0.826 + 1.615×1.118 + 1.975×1.419
            = 1.037 + 1.806 + 2.802
            = 5.645
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16

完整矩阵:

Q · K^T =
        位置0(猫)   位置1(追)   位置2(老鼠)
位置0   [ 7.212,    6.268,    5.645 ]    # 猫对所有词的关注
位置1   [ 6.135,    5.347,    4.813 ]    # 追对所有词的关注
位置2   [ 5.468,    4.789,    4.313 ]    # 老鼠对所有词的关注
1
2
3
4
5

3.2 为什么除以 √d_k? ​

问题:当 d_k 较大时,点积的值范围也较大,可能导致 softmax 饱和。

d_k = 3
√d_k = √3 ≈ 1.732
1
2

归一化:

原始 QK^T:
        位置0      位置1      位置2
位置0 [ 7.212,    6.268,    5.645 ]
位置1 [ 6.135,    5.347,    4.813 ]
位置2 [ 5.468,    4.789,    4.313 ]

除以 √d_k = 1.732 后:
        位置0      位置1      位置2
位置0 [ 4.165,    3.619,    3.260 ]
位置1 [ 3.543,    3.088,    2.779 ]
位置2 [ 3.157,    2.765,    2.491 ]
1
2
3
4
5
6
7
8
9
10
11

直观理解:

如果 d_k = 64,√d_k = 8
假设 Q 和 K 的每个元素都是独立的标准正态分布:
- 每个元素的期望 = 0,方差 = 1
- 点积 = 3个元素的和,方差 = 3,方差开方 ≈ 1.73
- 但实际 d_k=64 时,点积方差 = 64,标准差 = 8

直接 softmax([8, 0, 0, ...]) ≈ [1, 0, 0, ...]  → 梯度≈0,无法学习!

除以 √d_k 后:
- 点积方差归一化为 1
- softmax 分布更平滑,梯度正常
1
2
3
4
5
6
7
8
9
10
11

第4章:Softmax — 注意力权重 ​

4.1 Softmax 公式 ​

Attention_Weight[i,j] = exp(score[i,j] / √d_k) / Σ_j exp(score[i,j] / √d_k)

即:对每一行做 softmax(沿 K 的方向,即 j 轴)
1
2
3

4.2 手算 Softmax(以第0行"猫"为例) ​

归一化后的分数(第0行):
[4.165, 3.619, 3.260]

步骤1:计算每个元素的 exp():
exp(4.165) = e^4.165 ≈ 64.37
exp(3.619) = e^3.619 ≈ 37.43
exp(3.260) = e^3.260 ≈ 26.00

步骤2:求和:
sum = 64.37 + 37.43 + 26.00 = 127.80

步骤3:归一化:
weight(猫→猫)   = 64.37 / 127.80 = 0.504
weight(猫→追)   = 37.43 / 127.80 = 0.293
weight(猫→老鼠) = 26.00 / 127.80 = 0.203

验证:0.504 + 0.293 + 0.203 = 1.000 ✓
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17

4.3 完整注意力权重矩阵 ​

对所有三行计算 softmax:

注意力权重矩阵 A (3 × 3):

        位置0(猫)  位置1(追)  位置2(老鼠)
位置0   [ 0.504,    0.293,    0.203 ]  # 猫的关注分布
位置1   [ 0.428,    0.337,    0.235 ]  # 追的关注分布
位置2   [ 0.386,    0.342,    0.272 ]  # 老鼠的关注分布

解释:
- "猫"这个词,50.4% 关注自己,29.3% 关注"追",20.3% 关注"老鼠"
- "追"这个词,42.8% 关注"猫",33.7% 关注自己,23.5% 关注"老鼠"
- "老鼠"这个词,38.6% 关注"猫",34.2% 关注"追",27.2% 关注自己
1
2
3
4
5
6
7
8
9
10
11

4.4 可视化理解 ​

猫 → [猫: 50.4% | 追: 29.3% | 老鼠: 20.3%]
追 → [猫: 42.8% | 追: 33.7% | 老鼠: 23.5%]
老鼠 → [猫: 38.6% | 追: 34.2% | 老鼠: 27.2%]

注意力权重热力图:
         猫    追   老鼠
猫     0.50  0.29  0.20
追     0.43  0.34  0.24
老鼠   0.39  0.34  0.27
1
2
3
4
5
6
7
8
9

第5章:加权求和 — Attention Output ​

5.1 公式 ​

Attention_Output = A · V

即:注意力权重矩阵 × Value 矩阵
1
2
3

5.2 手算(以"猫"的输出为例) ​

output_猫 = 0.504 × V_猫 + 0.293 × V_追 + 0.203 × V_老鼠

V_猫   = [1.450, 1.810, 2.170]
V_追   = [1.245, 1.555, 1.875]
V_老鼠 = [1.095, 1.380, 1.675]

output_猫[0] = 0.504×1.450 + 0.293×1.245 + 0.203×1.095
             = 0.731 + 0.365 + 0.222
             = 1.318

output_猫[1] = 0.504×1.810 + 0.293×1.555 + 0.203×1.380
             = 0.912 + 0.456 + 0.280
             = 1.648

output_猫[2] = 0.504×2.170 + 0.293×1.875 + 0.203×1.675
             = 1.094 + 0.549 + 0.340
             = 1.983

output_猫 = [1.318, 1.648, 1.983]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19

5.3 完整 Attention Output ​

output_追  = [1.279, 1.597, 1.913]
output_老鼠 = [1.252, 1.565, 1.883]

Attention Output 矩阵 Z (3 × 3):
[[1.318, 1.648, 1.983],
 [1.279, 1.597, 1.913],
 [1.252, 1.565, 1.883]]
1
2
3
4
5
6
7

这就是 单头 Self-Attention 的最终输出。

5.4 单头注意力的完整数据流 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                    单头 Self-Attention 完整前向传播                            │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│  输入矩阵 X_input (3×4)                                                      │
│      │                                                                       │
│      ├─→ · W_q (4×3) ──→ Q (3×3)                                          │
│      ├─→ · W_k (4×3) ──→ K (3×3)                                          │
│      └─→ · W_v (4×3) ──→ V (3×3)                                          │
│                                                                             │
│  Q · K^T (3×3)                                                               │
│      ↓                                                                       │
│  ÷ √d_k (3×3)                                                               │
│      ↓                                                                       │
│  Softmax (3×3) ──→ A:注意力权重矩阵                                         │
│      ↓                                                                       │
│  A · V (3×3) ──→ Z:注意力输出矩阵                                           │
│                                                                             │
│  Z = softmax(Q · K^T / √d_k) · V                                           │
│                                                                             │
└─────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21

第6章:Multi-Head Attention — 多头注意力 ​

6.1 单头的局限性 ​

单头只能学习一种"注意力模式"

例如 "猫 追 老鼠":
- 单头可能只能学到"相邻词的关系"
- 但实际上句子中同时存在多种关系:
  1. 主谓关系(猫→追)
  2. 语义相似(老鼠和猫同属动物)
  3. 动作对象(追的目标是老鼠)

单头必须用同一个注意力模式表达所有关系,这限制了表达能力。
1
2
3
4
5
6
7
8
9
10

6.2 多头的解决方案 ​

每个头独立计算自己的 Attention,每个头有独立的 W_q, W_k, W_v
不同头学习不同类型的注意力模式

例子:
- Head₁:关注"主语-动词"关系
- Head₂:关注"动作-对象"关系
- Head₃:关注"语义相似性"
1
2
3
4
5
6
7

6.3 多头的参数分配 ​

设定:
- d_model = 4     (原始词向量维度)
- num_heads = 2   (2个头)
- d_k = d_v = 2   (每个头的维度)

总参数量不变:d_model × d_k × num_heads = 4×2×2 = 16
与单头 d_model × d_k = 4×4 = 16 相同

每个头处理 2 维(d_k=2),2个头共4维(d_model=4)
拼接后维度仍然回到 d_model
1
2
3
4
5
6
7
8
9
10

6.4 多头的具体计算 ​

Step 1:两个头分别计算 Attention

Head₁ 的投影矩阵:
W_q1 (4×2) = [[0.1, 0.2],     W_k1 (4×2) = [[0.05, 0.15],
              [0.4, 0.5],                   [0.35, 0.45],
              [0.7, 0.8],                   [0.65, 0.75],
              [0.15, 0.25]]                 [0.1, 0.2]]
W_v1 (4×2) = [[0.2, 0.3],
              [0.5, 0.6],
              [0.8, 0.9],
              [0.1, 0.2]]

Head₂ 的投影矩阵:
W_q2 (4×2) = [[0.3, 0.4],     W_k2 (4×2) = [[0.1, 0.2],
              [0.5, 0.6],                   [0.4, 0.5],
              [0.8, 0.9],                   [0.7, 0.8],
              [0.2, 0.3]]                   [0.15, 0.25]]
W_v2 (4×2) = [[0.3, 0.4],
              [0.6, 0.7],
              [0.9, 1.0],
              [0.2, 0.3]]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19

Step 2:各自计算 Attention(简化手算)

假设 X_input = [[0.50, 1.20, 0.80, 1.10],
                [1.14, 1.24, 0.20, 1.60],
                [1.51,-0.32, 0.30, 1.90]]

Head₁(d_k=2):
Q₁ = X · W_q1  → 3×2 矩阵
K₁ = X · W_k1  → 3×2 矩阵
V₁ = X · W_v1  → 3×2 矩阵
Z₁ = softmax(Q₁·K₁^T/√2) · V₁  → 3×2 矩阵

Head₂(d_k=2):
Q₂ = X · W_q2  → 3×2 矩阵
K₂ = X · W_k2  → 3×2 矩阵
V₂ = X · W_v2  → 3×2 矩阵
Z₂ = softmax(Q₂·K₂^T/√2) · V₂  → 3×2 矩阵
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15

简化假设(让手算可追踪):

假设经过计算后:

Z₁ (Head₁ 输出,3×2) = [[1.32, 1.65],
                        [1.28, 1.60],
                        [1.25, 1.57]]

Z₂ (Head₂ 输出,3×2) = [[1.40, 1.75],
                        [1.35, 1.68],
                        [1.30, 1.62]]
1
2
3
4
5
6
7
8
9

6.5 拼接(Concat) ​

所有头的输出沿维度方向拼接:
Z_concat = concat(Z₁, Z₂)  →  形状 (3, 2+2) = (3, 4)

        concat 沿 dim=1
Z₁: [[1.32, 1.65]]
Z₂: [[1.40, 1.75]]  ──→  [[1.32, 1.65, 1.40, 1.75]]  ← 位置0(猫)
                    ──→  [[1.28, 1.60, 1.35, 1.68]]  ← 位置1(追)
                    ──→  [[1.25, 1.57, 1.30, 1.62]]  ← 位置2(老鼠)

Z_concat (3 × 4):
[[1.32, 1.65, 1.40, 1.75],   # 猫
 [1.28, 1.60, 1.35, 1.68],   # 追
 [1.25, 1.57, 1.30, 1.62]]   # 老鼠
1
2
3
4
5
6
7
8
9
10
11
12
13

6.6 最终投影 W_o ​

Z_concat · W_o = MultiHead Attention 输出

W_o (4 × 4):
[[0.1, 0.2, 0.3, 0.4],
 [0.5, 0.6, 0.7, 0.8],
 [0.15, 0.25, 0.35, 0.45],
 [0.55, 0.65, 0.75, 0.85]]

MultiHead_output = Z_concat · W_o

示例计算(位置0的第0个元素):
output[0,0] = 1.32×0.1 + 1.65×0.5 + 1.40×0.15 + 1.75×0.55
            = 0.132 + 0.825 + 0.210 + 0.963
            = 2.130

最终 MultiHead Output (3 × 4):
[[2.130, 2.665, 3.195, 3.728],
 [2.066, 2.587, 3.103, 3.622],
 [2.018, 2.530, 3.037, 3.548]]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19

6.7 多头注意力全景图 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                    Multi-Head Attention 完整流程                             │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│                         X_input (3 × 4)                                     │
│                              │                                               │
│     ┌────────────────────────┼────────────────────────┐                     │
│     │                        │                        │                     │
│     ↓                        ↓                        ↓                     │
│  Head₁                   Head₂                   ... Head_h                │
│  W_q1, W_k1, W_v1    W_q2, W_k2, W_v2         W_qh, W_kh, W_vh           │
│     │                        │                        │                     │
│     ↓                        ↓                        ↓                     │
│  Z₁ (3×2)                Z₂ (3×2)                Zₕ (3×2)                 │
│     │                        │                        │                     │
│     └────────────────────────┼────────────────────────┘                     │
│                              ↓                                               │
│                    concat(Z₁, Z₂, ..., Zₕ)                                  │
│                              │                                               │
│                              ↓                                               │
│                    · W_o (4×4)                                              │
│                              │                                               │
│                              ↓                                               │
│               MultiHead Output (3 × 4)                                       │
│                                                                             │
│  MultiHead = concat(head₁, ..., headₕ) · W_o                               │
│                                                                             │
└─────────────────────────────────────────────────────────────────────────────┘
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

6.8 为什么多头有效? ​

类比:

单头 = 用一台显微镜观察切片(只能看到一种细节)
多头 = 同时用多台不同倍率的显微镜(看到不同尺度的信息)

技术原因:

1. 每个头独立学习不同的注意力模式
   - 某些头学习句法关系(主谓宾)
   - 某些头学习语义相似性
   - 某些头学习位置邻近性

2. 维度压缩不丢失信息
   - 2个头 × 2维 = 4维 = d_model
   - 拼接后维度恰好恢复

3. 冗余与鲁棒性
   - 某些头可能学习相似的模式(备份)
   - 某些头可能被"淘汰"(不重要)
   - 整体系统更健壮

实际观察(来自论文分析):
- 不同头确实学习了不同的语义关系
- 某些头专注于特定的语法结构
- 有些头对删除不敏感(冗余)
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

第7章:解码策略 — 从 Logits 到 Token ​

7.1 从向量到词语的完整链路 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                        Token 生成完整链路                                     │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│  MultiHead Output Z (3 × 4)                                                 │
│       │                                                                     │
│       ↓                                                                     │
│  LayerNorm + Add(残差连接)                                                 │
│       │                                                                     │
│       ↓                                                                     │
│  Feed-Forward Network                                                       │
│       │                                                                     │
│       ↓                                                                     │
│  LayerNorm + Add                                                            │
│       │                                                                     │
│       ↓                                                                     │
│  Linear (映射到词表)                                                          │
│       ↓                                                                     │
│  Logits (3 × 6)  ← 每个位置6个候选词的原始分数                               │
│       │                                                                     │
│       ↓                                                                     │
│  (训练时)Cross-Entropy Loss ← 计算与目标词的距离                            │
│  (推理时)Temperature → Softmax → TopK/TopP → 采样 → Token                  │
│                                                                             │
└─────────────────────────────────────────────────────────────────────────────┘
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

7.2 Logits 矩阵 ​

经过多层 Transformer 处理后,最后一层 Linear 将向量映射回词表大小:

Logits = Z · W_linear

假设词表大小 d_vocab = 6

Logits (3 × 6):
       <PAD>  老鼠   猫    追    我    爱
位置0 [ 0.1,  -0.5,  2.3,  1.2,  0.8, -0.3]   # 猫的位置
位置1 [ 0.2,   0.9, -0.2,  2.1,  0.5,  1.5]   # 追的位置
位置2 [-0.1,   2.5,  0.3,  0.7, -0.4,  1.8]   # 老鼠的位置

每个位置的 6 个值代表"选择该词"的原始得分(越高越可能)
1
2
3
4
5
6
7
8
9
10
11

7.3 从 Logits 到概率 — Softmax ​

P(token_i) = exp(logits_i) / Σ exp(logits_j)
1

以位置0(猫)为例:

位置0的 Logits:[0.1, -0.5, 2.3, 1.2, 0.8, -0.3]

exp计算:
exp(0.1)  = 1.105
exp(-0.5) = 0.607
exp(2.3)  = 9.974
exp(1.2)  = 3.320
exp(0.8)  = 2.226
exp(-0.3) = 0.741

总和 = 1.105 + 0.607 + 9.974 + 3.320 + 2.226 + 0.741 = 17.973

概率分布:
P(<PAD>) = 1.105/17.973 = 6.1%
P(老鼠)   = 0.607/17.973 = 3.4%
P(猫)    = 9.974/17.973 = 55.5%  ← 最高!
P(追)    = 3.320/17.973 = 18.5%
P(我)    = 2.226/17.973 = 12.4%
P(爱)    = 0.741/17.973 = 4.1%

总和 = 100% ✓
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21

第8章:Temperature — 控制随机性 ​

8.1 Temperature 的本质 ​

Temperature 是在 Softmax 之前对 Logits 进行缩放:

P(x) = softmax(logits / T) = exp(logits_i / T) / Σ exp(logits_j / T)

其中 T 是 Temperature 参数
1
2
3

8.2 数值对比:不同 Temperature 的效果 ​

以位置0为例,比较 T = 0.3, 0.7, 1.0, 1.5 的效果:

原始 Logits = [0.1, -0.5, 2.3, 1.2, 0.8, -0.3]

Step 1:除以 T
Step 2:exp
Step 3:归一化

━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
T = 0.3(低温 → 高确定性)
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
logits/T = [0.33, -1.67, 7.67, 4.00, 2.67, -1.00]
exp      = [1.39,  0.19, 2143,  54.6,  14.4,  0.368]
概率     = [0.06%, 0.01%, 94.3%, 2.4%, 0.6%,  0.02%]

→ "猫"占了94.3%,几乎必定选它(贪婪)

━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
T = 0.7(适中)
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
logits/T = [0.14, -0.71, 3.29, 1.71, 1.14, -0.43]
exp      = [1.15,  0.49, 26.8,  5.53,  3.13,  0.65]
概率     = [2.8%, 1.2%, 65.2%, 13.5%, 7.6%,  1.6%]

→ "猫"仍占主导(65%),但其他词有一定机会

━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
T = 1.0(原始分布)
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
概率 = [6.1%, 3.4%, 55.5%, 18.5%, 12.4%, 4.1%]

→ 这是未经 Temperature 调整的原始 Softmax

━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
T = 1.5(高温 → 高随机性)
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
logits/T = [0.07, -0.33, 1.53, 0.80, 0.53, -0.20]
exp      = [1.07,  0.72,  4.62, 2.23, 1.70,  0.82]
概率     = [8.0%, 5.4%, 34.5%, 16.6%, 12.7%, 6.1%]

→ "猫"从55%降到34%,其他词的机会增加
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━
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

8.3 Temperature 的直观理解 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                    Temperature 如何"塑造"概率分布                              │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│  T → 0(极低):  ████████████████████████████████████ "猫" 94%             │
│                  ██                                  "追"  2%               │
│                  ░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 其他 4%               │
│                  → 几乎总是选最高概率,类似贪婪解码                             │
│                                                                             │
│  T = 1.0(原始):███████ "猫" 55%                                           │
│                  ████       "追" 19%                                        │
│                  ██         "我" 12%                                        │
│                  █          其他 14%                                        │
│                  → 保持模型学习到的原始分布                                    │
│                                                                             │
│  T → ∞(极高):  ██████████ "猫" 35%                                        │
│                  ████████   "追" 25%                                        │
│                  ██████     "我" 20%                                        │
│                  █████      其他 20%                                        │
│                  → 分布被"拉平",接近均匀随机                                  │
│                                                                             │
└─────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22

第9章:TopK 采样 — 只选最好的 K 个 ​

9.1 TopK 的原理 ​

TopK = 限制候选 Token 数量,只考虑概率最高的 K 个:

步骤:
1. 将所有 Token 按概率从高到低排序
2. 只保留前 K 个,其余概率设为 0
3. 重新归一化概率
4. 从这 K 个中采样
1
2
3
4
5

9.2 具体示例 ​

以位置0为例,假设 d_vocab = 6,设定 TopK = 3:

原始概率分布(已排序):
Token:  猫     追     我     <PAD>  老鼠   爱
Prob:  0.555  0.185  0.124  0.061  0.034  0.041  ← 注意老鼠和爱交换了位置

TopK = 3,选择前3个:[猫, 追, 我]

新分布:
Token:  猫     追     我     <PAD>  老鼠   爱
Prob:  0.555  0.185  0.124  0.061  0.034  0.041
       ────────────────────────────  ─────────────────
       参与采样                        设为0

重新归一化(除以 0.555+0.185+0.124 = 0.864):
Token:  猫     追     我     <PAD>  老鼠   爱
Prob:  0.642  0.214  0.144  0.000  0.000  0.000

最终采样范围:
- "猫":64.2%
- "追":21.4%
- "我":14.4%
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

9.3 TopK 的优缺点 ​

优点:
✓ 有效过滤极低概率的噪声 Token
✓ 采样数量固定,计算稳定
✓ 适合需要一定多样性的场景

缺点:
✗ K 是固定值,在不同位置可能不合适
  - 当前一分布均匀时,TopK=3 可能仍然包含很多低概率词
  - 当前一很集中时,TopK=3 可能只保留了真正的高概率词

TopK 场景选择:
- K=1:贪婪解码(退化为取最高概率)
- K=10~50:一般对话
- K=50~100:需要高多样性的生成任务
1
2
3
4
5
6
7
8
9
10
11
12
13
14

第10章:TopP 采样 — 核采样(Nucleus Sampling) ​

10.1 TopP 的原理 ​

TopP = 动态选择累积概率达到阈值 P 的最小 Token 集合:

步骤:
1. 按概率从高到低排序
2. 从高到低累加概率
3. 当累积概率首次达到 P 时,停止
4. 只从这些 Token 中采样
1
2
3
4
5

10.2 具体示例 ​

原始概率分布(已排序):
Token:   猫      追      我     <PAD>   爱     老鼠
Prob:   0.555  0.185  0.124  0.061  0.041  0.034
累计:   0.555  0.740  0.864  0.925  0.966  1.000

TopP = 0.9(累积90%):
从高到低累加:
- "猫":0.555 → 累计0.555,不够0.9,继续
- "追":0.185 → 累计0.740,不够0.9,继续
- "我":0.124 → 累计0.864,不够0.9,继续
- "<PAD>":0.061 → 累计0.925,超过0.9,停止!

候选集合:{猫, 追, 我, <PAD>}  ← 这4个词贡献了90%以上的概率

重新归一化这4个词:
总和 = 0.555 + 0.185 + 0.124 + 0.061 = 0.925

新概率:
Token:   猫     追     我    <PAD>   爱    老鼠
Prob:   0.600  0.200  0.134  0.066  0.000  0.000
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

10.3 TopP 与 TopK 对比 ​

TopK = 固定数量,动态范围
┌─────────────────────────────────────────────┐
│ Token  概率    TopK=3                       │
├─────────────────────────────────────────────┤
│ 猫     0.555  ✓ 参与                        │
│ 追     0.185  ✓ 参与                        │
│ 我     0.124  ✓ 参与                        │
│ <PAD>  0.061  ✗ 被排除                      │
│ 爱     0.041  ✗ 被排除                      │
│ 老鼠   0.034  ✗ 被排除                      │
└─────────────────────────────────────────────┘

TopP = 动态数量,固定概率质量
┌─────────────────────────────────────────────┐
│ Token  概率    累计    TopP=0.9             │
├─────────────────────────────────────────────┤
│ 猫     0.555  0.555   ✓                    │
│ 追     0.185  0.740   ✓                    │
│ 我     0.124  0.864   ✓                    │
│ <PAD>  0.061  0.925   ✓ (首次超过0.9)       │
│ 爱     0.041  0.966   ✗                     │
│ 老鼠   0.034  1.000   ✗                     │
└─────────────────────────────────────────────┘

TopP 选了4个词,TopK 选了3个词。
哪个更合理?取决于场景!
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

10.4 TopP 的优势 ​

为什么现代 LLM 默认推荐 TopP 而非 TopK?

1. 动态适应分布形状
   - 分布集中时 → 自动选择少量 Token
   - 分布分散时 → 自动扩大 Token 范围

2. 避免极端情况
   - TopK=3 但如果前3个词概率都很低(各10%),实际只覆盖了30%
   - TopP=0.9 会继续往下选,直到覆盖90%

3. 实践效果好
   - OpenAI、Anthropic 等厂商的 API 默认参数都使用 TopP
   - 适合大多数场景,无需手动调参
1
2
3
4
5
6
7
8
9
10
11
12
13

第11章:综合对比与实践指南 ​

11.1 Temperature vs TopK vs TopP ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                          三种策略对比                                         │
├──────────────────┬──────────────────┬──────────────────┬────────────────────┤
│ 特性             │ Temperature      │ TopK             │ TopP               │
├──────────────────┼──────────────────┼──────────────────┼────────────────────┤
│ 作用对象         │ Logits(分数)    │ 概率排序后        │ 累积概率           │
│ 控制方式         │ 缩放分布形状      │ 截断候选数量      │ 截断概率质量       │
│ 常见值           │ 0.0 ~ 2.0        │ 1 ~ 100          │ 0.0 ~ 1.0          │
│ 默认值           │ 1.0              │ 1 (贪婪)         │ 0.9 / 0.95         │
│ 类比             │ 让差异更大/更小   │ 只看最好的几个    │ 覆盖绝大多概率     │
│ 典型应用         │ 控制确定性程度    │ 平衡多样+质量    │ 推理时推荐         │
└──────────────────┴──────────────────┴──────────────────┴────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12

11.2 场景化推荐参数 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                        场景化解码参数推荐                                     │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│  场景            │ Temperature │ TopP  │ TopK │ 说明                        │
│ ─────────────────┼─────────────┼───────┼──────┼─────────────────────────── │
│  代码生成         │   0.0~0.3   │ 1.0   │  1   │ 确定性强,语法必须正确        │
│  数学证明         │   0.0~0.2   │ 1.0   │  1   │ 精确答案,不可随机            │
│  正式邮件         │   0.3~0.5   │ 0.9   │  50  │ 专业稳定,轻微变化            │
│  问答系统         │   0.5~0.7   │ 0.9   │  40  │ 平衡准确与自然                │
│  日常对话         │   0.7~0.9   │ 0.95  │  50  │ 自然流畅,适度随机            │
│  创意写作         │   0.9~1.2   │ 0.95  │ 100  │ 高多样性,鼓励创新            │
│  头脑风暴         │   1.0~1.5   │ 0.8   │ 200  │ 最高多样性,不惜牺牲连贯性    │
│  角色扮演         │   0.8~1.0   │ 0.95  │  80  │ 符合人物性格的同时有变化      │
│                                                                             │
└─────────────────────────────────────────────────────────────────────────────┘
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16

11.3 采样完整流程图 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                      推理时 Token 生成完整流程                                 │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│  1. 输入 Token 序列                                                           │
│       ↓                                                                       │
│  2. Embedding + 位置编码 → X_input                                          │
│       ↓                                                                       │
│  3. Transformer 编码器层(×N)                                               │
│       ↓                                                                       │
│  4. Transformer 解码器层(×N)                                                │
│       ↓                                                                       │
│  5. Linear 投影 → Logits (1 × d_vocab)                                     │
│       ↓                                                                       │
│  ┌─────────────────────────────────────────────────────────────┐           │
│  │ 解码策略应用(以 TopP 为例):                                  │           │
│  │                                                              │           │
│  │ Logits = [0.1, -0.5, 2.3, 1.2, 0.8, -0.3]                   │           │
│  │      ↓                                                        │           │
│  │ (可选)÷ Temperature                                          │           │
│  │      ↓                                                        │           │
│  │ Softmax → 概率分布                                             │           │
│  │      ↓                                                        │           │
│  │ 按概率排序 + 计算累积概率                                        │           │
│  │      ↓                                                        │           │
│  │ TopP 截断 → 候选 Token 集合                                    │           │
│  │      ↓                                                        │           │
│  │ 归一化候选概率                                                  │           │
│  │      ↓                                                        │           │
│  │ 按概率采样 → 选中 Token                                        │           │
│  └─────────────────────────────────────────────────────────────┘           │
│       ↓                                                                       │
│  6. 输出 Token(如 "爱")                                                     │
│       ↓                                                                       │
│  7. 将 "爱" 加入序列,重复步骤 1~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

第12章:训练阶段 vs 推理阶段 — 对比总结 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                        训练阶段 vs 推理阶段                                   │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│  【训练阶段】                                                                 │
│  ┌───────────────────────────────────────────────────────────────────────┐  │
│  │ 输入序列:"猫 追 老鼠 <EOS>"                                           │  │
│  │                                                                       │  │
│  │ Forward Pass:                                                       │  │
│  │   X_input → MultiHead → FFN → LayerNorm → Linear → Logits           │  │
│  │                                                                       │  │
│  │ Loss 计算:                                                           │  │
│  │   对每个位置,用 Logits 计算与"正确答案"(标签)的 Cross-Entropy Loss   │  │
│  │                                                                       │  │
│  │   position0: 正确答案="追"   → loss₀                                 │  │
│  │   position1: 正确答案="老鼠" → loss₁                                 │  │
│  │   position2: 正确答案="<EOS>" → loss₂                                │  │
│  │                                                                       │  │
│  │   Total Loss = (loss₀ + loss₁ + loss₂) / 3                           │  │
│  │                                                                       │  │
│  │ Backward Pass:                                                       │  │
│  │   Loss → 反向传播梯度 → 更新 W_q, W_k, W_v, W_o, FFN, Embedding       │  │
│  │                                                                       │  │
│  │ 特点:                                                                │  │
│  │ ✓ 可以并行处理整个序列(Teacher Forcing)                               │  │
│  │ ✓ 所有位置的损失同时计算                                               │  │
│  │ ✓ 不使用 Temperature/TopP(直接用 Cross-Entropy)                      │  │
│  └───────────────────────────────────────────────────────────────────────┘  │
│                                                                             │
│  【推理阶段】                                                                 │
│  ┌───────────────────────────────────────────────────────────────────────┐  │
│  │ 输入序列:"猫"                                                         │  │
│  │                                                                       │  │
│  │ Forward Pass:                                                        │  │
│  │   X_input → MultiHead → FFN → LayerNorm → Linear → Logits            │  │
│  │                                                                       │  │
│  │ 解码策略:                                                             │  │
│  │   Logits → Temperature → Softmax → TopP → 采样 → Token               │  │
│  │                                                                       │  │
│  │ 特点:                                                                │  │
│  │ ✗ 自回归生成(AR),必须一个 Token 一个 Token 地生成                   │  │
│  │ ✗ 每个位置只能用已生成的前文                                            │  │
│  │ ✗ 需要解码策略控制输出                                                 │  │
│  └───────────────────────────────────────────────────────────────────────┘  │
│                                                                             │
└─────────────────────────────────────────────────────────────────────────────┘
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

核心公式速查卡 ​

┌─────────────────────────────────────────────────────────────────────────────┐
│                          Transformer 核心公式速查                             │
├─────────────────────────────────────────────────────────────────────────────┤
│                                                                             │
│  【输入】                                                                    │
│  X_input = Embedding(Token) + Positional_Encoding(pos)                     │
│                                                                             │
│  【QKV 投影】                                                                │
│  Q = X · W_q                                                                │
│  K = X · W_k                                                                │
│  V = X · W_v                                                                │
│                                                                             │
│  【单头 Attention】                                                          │
│  Attention(Q, K, V) = softmax(Q · K^T / √d_k) · V                          │
│                                                                             │
│  【多头 Attention】                                                          │
│  MultiHead = concat(head₁, ..., headₕ) · W_o                               │
│  where headᵢ = Attention(X · W_qᵢ, X · W_kᵢ, X · W_vᵢ)                    │
│                                                                             │
│  【位置编码】                                                                │
│  PE(pos, 2i)   = sin(pos / 10000^(2i/d))                                   │
│  PE(pos, 2i+1) = cos(pos / 10000^(2i/d))                                   │
│                                                                             │
│  【解码策略】                                                                │
│  Temperature: P = softmax(logits / T)                                       │
│  TopK:          只保留概率最高的 K 个 Token                                   │
│  TopP:          保留累积概率达到 P 的最小 Token 集合                          │
│                                                                             │
└─────────────────────────────────────────────────────────────────────────────┘
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

章节测试 ​

测试1:QKV 计算 ​

假设输入向量 x = [1, 0, 1, 0],W_q = [[1,0],[0,1],[1,0],[0,1]],计算 Q 向量。

测试2:注意力权重 ​

假设两个词的注意力分数为 [3.0, 1.0],求 softmax 后的注意力权重。

测试3:多头拼接 ​

如果有两个头,Head₁ 输出 Z₁(2×3),Head₂ 输出 Z₂(2×3),concat 后的矩阵形状是什么?

测试4:Temperature 效果 ​

当 Temperature = 0.5 时,一个高概率词(logits=2.0)和低概率词(logits=-1.0)的概率差距会变大还是变小?

测试5:TopP 采样 ​

概率分布 [0.50, 0.25, 0.15, 0.10],TopP=0.8 会选中哪几个 Token?


参考答案 ​

测试1答案 ​

Q = x · W_q
Q[0] = 1×1 + 0×0 + 1×1 + 0×0 = 2
Q[1] = 1×0 + 0×1 + 1×0 + 0×1 = 0
Q = [2, 0]
1
2
3
4

测试2答案 ​

exp(3.0) = 20.09
exp(1.0) = 2.72
总和 = 22.81

weight₁ = 20.09 / 22.81 ≈ 0.881
weight₂ = 2.72 / 22.81 ≈ 0.119
1
2
3
4
5
6

测试3答案 ​

(2, 3+3) = (2, 6)

测试4答案 ​

变大。T < 1 时,softmax 趋向于放大高概率、压低低概率之间的差距。

测试5答案 ​

累积:0.50 → 0.75 → 0.90(首次超过0.9),所以选中前3个:[0.50, 0.25, 0.15]


相关笔记 ​

  • [[01 - 神经网络基础]] — 矩阵运算、激活函数的数学基础
  • [[07 - LLM 进化史]] — Transformer 的历史背景与革命性意义
  • [[08 - Transformer 原理]] — 架构层面的理解(组件视角)
  • [[03 - 解码策略]] — 从应用角度的参数调优指南

学习状态:🟡 新建

最后更新于:

Pager
上一篇12. 训练 vs 推理:同一个 Transformer,两条完全不同的执行路径
下一篇14. Transformer 推理阶段详解 — 模型如何"思考"并生成回答 / Transformer Inference and Autoregressive Generation

持续记录,持续成长

Copyright © Tidenflow