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)为了手算方便,我们约定以下超参数:
d_model = 4 # 词向量维度(最终输出向量维度)
seq_len = 3 # 序列长度(3个词)
d_vocab = 6 # 假设词表大小为61.2 Token 到索引
首先,每个词被映射为一个整数索引(词表中的位置):
词表(按字母排序,仅作示例):
索引 0: "<PAD>" (填充符)
索引 1: "老鼠"
索引 2: "猫"
索引 3: "追"
索引 4: "我"
索引 5: "爱"
输入序列 "猫 追 老鼠" 对应的索引:
Token 索引
--------------
猫 → 2
追 → 3
老鼠 → 11.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] # 爱通过索引查找,得到每个词的初始向量:
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.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]具体计算(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]位置编码矩阵 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.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]]这就是传入 Transformer 层的第一个输入。
第2章:QKV 矩阵 — 线性投影
2.1 Q/K/V 的核心思想
Query(查询): 我当前这个词想要查找什么信息?
Key(键) : 序列中每个位置能提供什么信息?
Value(值) : 每个位置的实际信息内容这三个向量都是通过同一个输入向量 X_input 做不同的线性投影得到的:
Q = X_input · W_q (Query 投影)
K = X_input · W_k (Key 投影)
V = X_input · W_v (Value 投影)投影矩阵 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]]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:
位置0 (猫): Q_猫 = [1.255, 1.615, 1.975]
位置1 (追): Q_追 = [1.088, 1.404, 1.732] # 略去中间计算
位置2 (老鼠): Q_老鼠 = [0.967, 1.269, 1.582] # 略去中间计算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] # 老鼠的查询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.795K 矩阵 (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] # 老鼠的键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.17V 矩阵 (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] # 老鼠的值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]]第3章:注意力分数计算 — QK^T / √d_k
3.1 为什么是 Q · K^T?
Q 问"我需要什么信息",K 说"我能提供什么信息"。
两者的点积衡量了"当前词对序列中其他词的关注程度":
score(i, j) = Q[i] · K[j]^T = 第i个词的查询 与 第j个词的键 的相似度计算 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完整矩阵:
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 ] # 老鼠对所有词的关注3.2 为什么除以 √d_k?
问题:当 d_k 较大时,点积的值范围也较大,可能导致 softmax 饱和。
d_k = 3
√d_k = √3 ≈ 1.732归一化:
原始 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 ]直观理解:
如果 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 分布更平滑,梯度正常第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 轴)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 ✓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% 关注自己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第5章:加权求和 — Attention Output
5.1 公式
Attention_Output = A · V
即:注意力权重矩阵 × Value 矩阵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]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]]这就是 单头 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 │
│ │
└─────────────────────────────────────────────────────────────────────────────┘第6章:Multi-Head Attention — 多头注意力
6.1 单头的局限性
单头只能学习一种"注意力模式"
例如 "猫 追 老鼠":
- 单头可能只能学到"相邻词的关系"
- 但实际上句子中同时存在多种关系:
1. 主谓关系(猫→追)
2. 语义相似(老鼠和猫同属动物)
3. 动作对象(追的目标是老鼠)
单头必须用同一个注意力模式表达所有关系,这限制了表达能力。6.2 多头的解决方案
每个头独立计算自己的 Attention,每个头有独立的 W_q, W_k, W_v
不同头学习不同类型的注意力模式
例子:
- Head₁:关注"主语-动词"关系
- Head₂:关注"动作-对象"关系
- Head₃:关注"语义相似性"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_model6.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]]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 矩阵简化假设(让手算可追踪):
假设经过计算后:
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]]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]] # 老鼠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]]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 │
│ │
└─────────────────────────────────────────────────────────────────────────────┘6.8 为什么多头有效?
类比:
单头 = 用一台显微镜观察切片(只能看到一种细节)
多头 = 同时用多台不同倍率的显微镜(看到不同尺度的信息)
技术原因:
1. 每个头独立学习不同的注意力模式
- 某些头学习句法关系(主谓宾)
- 某些头学习语义相似性
- 某些头学习位置邻近性
2. 维度压缩不丢失信息
- 2个头 × 2维 = 4维 = d_model
- 拼接后维度恰好恢复
3. 冗余与鲁棒性
- 某些头可能学习相似的模式(备份)
- 某些头可能被"淘汰"(不重要)
- 整体系统更健壮
实际观察(来自论文分析):
- 不同头确实学习了不同的语义关系
- 某些头专注于特定的语法结构
- 有些头对删除不敏感(冗余)第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 │
│ │
└─────────────────────────────────────────────────────────────────────────────┘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 个值代表"选择该词"的原始得分(越高越可能)7.3 从 Logits 到概率 — Softmax
P(token_i) = exp(logits_i) / Σ exp(logits_j)以位置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% ✓第8章:Temperature — 控制随机性
8.1 Temperature 的本质
Temperature 是在 Softmax 之前对 Logits 进行缩放:
P(x) = softmax(logits / T) = exp(logits_i / T) / Σ exp(logits_j / T)
其中 T 是 Temperature 参数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%,其他词的机会增加
━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━━8.3 Temperature 的直观理解
┌─────────────────────────────────────────────────────────────────────────────┐
│ Temperature 如何"塑造"概率分布 │
├─────────────────────────────────────────────────────────────────────────────┤
│ │
│ T → 0(极低): ████████████████████████████████████ "猫" 94% │
│ ██ "追" 2% │
│ ░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░░ 其他 4% │
│ → 几乎总是选最高概率,类似贪婪解码 │
│ │
│ T = 1.0(原始):███████ "猫" 55% │
│ ████ "追" 19% │
│ ██ "我" 12% │
│ █ 其他 14% │
│ → 保持模型学习到的原始分布 │
│ │
│ T → ∞(极高): ██████████ "猫" 35% │
│ ████████ "追" 25% │
│ ██████ "我" 20% │
│ █████ 其他 20% │
│ → 分布被"拉平",接近均匀随机 │
│ │
└─────────────────────────────────────────────────────────────────────────────┘第9章:TopK 采样 — 只选最好的 K 个
9.1 TopK 的原理
TopK = 限制候选 Token 数量,只考虑概率最高的 K 个:
步骤:
1. 将所有 Token 按概率从高到低排序
2. 只保留前 K 个,其余概率设为 0
3. 重新归一化概率
4. 从这 K 个中采样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%9.3 TopK 的优缺点
优点:
✓ 有效过滤极低概率的噪声 Token
✓ 采样数量固定,计算稳定
✓ 适合需要一定多样性的场景
缺点:
✗ K 是固定值,在不同位置可能不合适
- 当前一分布均匀时,TopK=3 可能仍然包含很多低概率词
- 当前一很集中时,TopK=3 可能只保留了真正的高概率词
TopK 场景选择:
- K=1:贪婪解码(退化为取最高概率)
- K=10~50:一般对话
- K=50~100:需要高多样性的生成任务第10章:TopP 采样 — 核采样(Nucleus Sampling)
10.1 TopP 的原理
TopP = 动态选择累积概率达到阈值 P 的最小 Token 集合:
步骤:
1. 按概率从高到低排序
2. 从高到低累加概率
3. 当累积概率首次达到 P 时,停止
4. 只从这些 Token 中采样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.00010.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个词。
哪个更合理?取决于场景!10.4 TopP 的优势
为什么现代 LLM 默认推荐 TopP 而非 TopK?
1. 动态适应分布形状
- 分布集中时 → 自动选择少量 Token
- 分布分散时 → 自动扩大 Token 范围
2. 避免极端情况
- TopK=3 但如果前3个词概率都很低(各10%),实际只覆盖了30%
- TopP=0.9 会继续往下选,直到覆盖90%
3. 实践效果好
- OpenAI、Anthropic 等厂商的 API 默认参数都使用 TopP
- 适合大多数场景,无需手动调参第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 │
│ 类比 │ 让差异更大/更小 │ 只看最好的几个 │ 覆盖绝大多概率 │
│ 典型应用 │ 控制确定性程度 │ 平衡多样+质量 │ 推理时推荐 │
└──────────────────┴──────────────────┴──────────────────┴────────────────────┘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 │ 符合人物性格的同时有变化 │
│ │
└─────────────────────────────────────────────────────────────────────────────┘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 │
│ │
└─────────────────────────────────────────────────────────────────────────────┘第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 地生成 │ │
│ │ ✗ 每个位置只能用已生成的前文 │ │
│ │ ✗ 需要解码策略控制输出 │ │
│ └───────────────────────────────────────────────────────────────────────┘ │
│ │
└─────────────────────────────────────────────────────────────────────────────┘核心公式速查卡
┌─────────────────────────────────────────────────────────────────────────────┐
│ 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: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]测试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测试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 - 解码策略]] — 从应用角度的参数调优指南
学习状态:🟡 新建