深度学习进阶 · 从原理到研究方法
这套材料的定位:把「会写代码」变成「能判断」。
上一套 PyTorch 入门教你怎么做;这一套告诉你为什么这样做、什么时候不该这样做。
目标是让你能独立读论文、独立判断一个方案对不对、自己找出改进点。
这套材料和网上教程的区别
网上(包括大部分系统课)的问题是:告诉你「用残差连接能解决退化」,但不告诉你为什么能、什么条件下会失效、怎么判断一个方案是真的有效还是在过拟合。
这套材料的写法是:
- 推导优先。每个概念都用数学推导讲一遍,不用「可以理解为」糊弄。
- 给出边界。每个方法都说明它在什么条件下成立、什么时候崩。
- 可验证。每篇有可运行的代码,验证公式确实成立。
数学要求:会用导数、会矩阵乘法、理解条件概率即可。深度学习里用到的数学比你想的浅——真正的难点在直觉和判断力,不在公式难度。
目录与学习顺序
左侧目录树就是完整结构,这里再给一份可以一次看完的总表。
Part 1 · 机器学习基础(4 篇)
这一部分是地基。跳过它,后面所有内容都是空中楼阁。算法岗笔试和论文里反复出现的概念都在这里。
| # | 文档 | 核心问题 |
|---|---|---|
| 1 | 机器学习的本质是什么 | 学习=搜索?损失函数到底在衡量什么? |
| 2 | 泛化、过拟合与偏差方差 | 训练集好测试集差,为什么? |
| 3 | 正则化与容量控制 | L1/L2 到底在做什么?为什么 weight decay 有时不等价? |
| 4 | 优化器:从 SGD 到 AdamW | Momentum/Adam 的更新量怎么推出来的?AdamW 修好了什么? |
Part 2 · 深度学习原理(5 篇)
这一部分讲深度网络特有的机制。为什么深网络能工作、为什么会出现梯度消失、以及那些"魔法数字"(为什么是 1e-3、为什么 warmup)背后是什么。
| # | 文档 | 核心问题 |
|---|---|---|
| 5 | 反向传播的完整推导 | 链式法则在这个网络里怎么具体展开?梯度消失的本质是什么? |
| 6 | 归一化:从 BatchNorm 到 RMSNorm | 归一化到底在解决什么?为什么 BatchNorm 淘汰而 LayerNorm 统治? |
| 7 | 卷积与感受野 | 卷积为什么有效?什么条件下等价于全连接? |
| 8 | 残差连接与网络架构 | 退化问题的数学本质是什么? |
| 9 | 初始化与数值稳定 | 为什么不能用 0 初始化?为什么深层需要特定初始化? |
Part 3 · Transformer 与现代 LLM(5 篇)
这一部分是当前主流架构。从 Attention 的推导一路到 2024-2025 年的主流设计。
| # | 文档 | 核心问题 |
|---|---|---|
| 10 | Attention 的数学推导 | QKV 为什么要除 √d?softmax 为什么必须有? |
| 11 | 多头注意力与位置编码 | 多头在多做什么?位置信息怎么注入?RoPE 的旋转思想 |
| 12 | 从 Transformer 到 LLaMA 系列 | 现代 LLM 改了什么?RMSNorm/SwiGLU/RoPE 各解决了什么? |
| 13 | 推理优化:KV Cache 与 GQA | 自回归生成为什么慢?缓存和 GQA 怎么解决? |
| 14 | 缩放定律与高效注意力 | 为什么可以「大力出奇迹」?FlashAttention 赢在哪? |
Part 4 · 研究方法论(2 篇)
这一部分是从「会学」到「会研究」。算法岗面试和实际研究,靠的是这套东西。
| # | 文档 | 核心问题 |
|---|---|---|
| 15 | 怎么读一篇论文 | 三遍法怎么用?如何判断一个方法的真实贡献? |
| 16 | 怎么设计实验 | 消融实验怎么做才可信?如何避免自欺欺人? |
Part 5 · 附录
| 文档 | 内容 |
|---|---|
| 公式速查表 | 所有关键公式一页汇总 |
| 面试高频问题 | 算法岗面试会问的 30 个问题(含答案要点) |
| llama_from_scratch.py | 从零实现的 LLaMA(RMSNorm + RoPE + GQA + SwiGLU + Pre-LN),已实测前向 + 反向通过 |
| 深度学习进阶手册.html | 单页离线阅读版(19 篇全部打包,含公式与代码高亮) |
三条学习路线
路线 A · 打基础 + 读论文(推荐,最扎实)
按 1 → 16 顺序完整走一遍。适合有 2-3 个月、目标是能独立读论文和做研究。
路线 B · 只要工程能力(最快落地)1 → 2 → 4 → 6 → 10 → 12 → 13。跳过推导细节,但 Part 3 全部要读——那是当前架构的必需品。
路线 C · 面试突击
直接看 面试高频问题,按题目反查对应文档。重点:Part 1 全部 + Part 2 的推导 + Part 3 的 Attention。
前置知识清单
开始前确认你会这些,不会就先补:
必须会
- 导数、偏导、链式法则
- 矩阵乘法、求逆(不用手算,知道意思即可)
- 概率:条件概率、贝叶斯、期望、方差
- Python + NumPy 基本操作
不用会(会查就行)
- 矩阵特征值、SVD 推导
- 大数定律的严格证明
- 卷积的傅里叶视角
边学边补
- 如果梯度公式不熟 → 先看 Part 2 第 5 篇,它会完整推导一遍
- 如果注意力机制第一次接触 → 直接从 Part 3 第 10 篇开始,那里从零推导
每篇的读法建议
不要只读。 每篇都有代码,建议流程:
1. 先读推导,假装自己懂了
2. 手推一遍关键公式(纸上)
3. 跑代码验证你推的结果
4. 合上文档,用自己的话复述「这个方法解决什么问题」
5. 做最后的「边界条件」自测题2
3
4
5
第 4 步是关键。如果你能讲清楚一个方法在什么情况下会失效,你才是真懂了。
关于代码
完整的 LLaMA 实现:llama_from_scratch.py,163 行,可运行:
pip install torch
python llama_from_scratch.py # 前向 + loss + 反向,全程通过2
它包含 RMSNorm、RoPE、GQA、SwiGLU、Pre-LN 全部现代组件,可以直接作为你的实现参考。
文档里每段代码的输出都对应某个公式或结论(都是实测的)。如果你推出来的数和代码算出来的不一样,先信代码。
从第 1 篇开始 → 机器学习的本质是什么
Part 1 · 机器学习基础
这一部分是地基。跳过它,后面所有内容都是空中楼阁 —— 算法岗笔试和论文里反复出现的概念都在这里。
四条线,按顺序看:
| # | 章节 | 核心问题 |
|---|---|---|
| 01 | 机器学习的本质是什么 | 学习 = 搜索?损失函数到底在衡量什么? |
| 02 | 泛化、过拟合与偏差方差 | 训练集好测试集差,为什么? |
| 03 | 正则化与容量控制 | L1/L2 到底在做什么?为什么 weight decay 有时不等价? |
| 04 | 优化器:从 SGD 到 AdamW | Momentum / Adam 的更新量怎么推出来的?AdamW 修好了什么? |
这一部分要建立的判断力
- 看到「模型效果好」时,先问:好在哪个集合上、和什么比
- 看到「加了正则化」时,先问:它约束的是什么,是参数范数还是函数复杂度
- 看到「换了个优化器」时,先问:它改的是步长还是方向
顺序不能跳
第 2 篇(偏差方差)是后面所有内容的公共语言。 第 5 篇讲梯度消失时会直接引用它,第 14 篇讲缩放定律也会。
1 · 机器学习的本质是什么
核心问题:模型在学什么?损失函数在衡量什么?为什么「最小化损失」等于「学得好」?
一、一个不准确的直觉
多数入门材料会说:「机器学习就是让计算机从数据中学习规律」。这句话没错,但没有任何信息量。
更有用的定义是:
机器学习 = 在一个函数空间里搜索,使得在某个目标函数上取值最小。
三个关键词,逐个拆。
关键词一:函数空间
你写的模型 f(x; θ) 里,θ 是参数。参数固定,模型就固定了。
假设模型是线性回归 f(x) = wx + b,那么 w 和 b 就是两个自由变量。它们构成的二维平面就是这个模型的函数空间。
| 假设 | 函数空间 |
|---|---|
| 线性回归(2 参数) | 二维平面 |
| 线性回归(10 参数) | 十维空间 |
| 一个 1 亿参数的 Transformer | 一亿维空间 |
机器学习的第一个关键事实:参数空间极其巨大,找到一个"好"的点,和找到一个"最优"的点,是两回事。
深度学习的全部困难都源于此——在一亿维空间里找最优解,实际做不到,只能不断改善。
关键词二:搜索
有了函数空间,接下来要决定「往哪个方向走」。这一步叫优化,由优化器(optimizer)完成。
在深度学习里,你写的这一行:
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)就是决定了搜索策略。而 loss.backward() 是在计算方向(往哪走能让损失变小)。
「方向 + 步长」构成了优化的全部。 后面所有优化器的区别,本质上只是这两件事的不同做法。
关键词三:目标函数
你优化的是 loss_fn(pred, y)。这个函数衡量的东西,决定了模型学到什么。
- 用
MSELoss→ 模型学「预测值和真实值的平方距离」 - 用
CrossEntropyLoss→ 模型学「预测类别的对数似然」 - 用对比学习损失(如 InfoNCE)→ 模型学「什么样本该靠近、什么该远离」
同一个模型架构,换个损失函数就是在做完全不同的事。 这是很多人换 loss 就想提升准确率却没效果的原因——他们换了,但没换对目标。
二、损失函数到底是什么
这一节是本文的核心。搞清楚损失函数的来源,后面所有选择都有依据。
从概率建模推导出来
假设你要做二分类。数据是 (x,y),其中 y∈{0,1}。
建模思路:假设存在一个真实概率 p(y=1∣x),我们的模型输出一个 p^。目标是让模型输出的概率尽可能接近真实概率。
用什么衡量两个概率分布的差异? KL 散度:
DKL(p∥q)=i∑pilogqipi
直观含义:如果真实分布是 p,我们用 q 去编码,编码平均长度比最优编码多多少比特。越接近 0 越好。
(推论:为什么 KL 散度不对称——DKL(p∥q)=DKL(q∥p)。这在后面的对比学习里会再次出现。)
从最大似然到交叉熵
KL 散度需要真实分布 p,但我们不知道。改用「最大似然」:找到一组参数,让观测到的数据出现的概率最大。
L=−i∑logp^(yi∣xi)
即负对数似然(NLL)。因为最大化 ∏p 等价于最小化 ∑−logp。
展开二分类(y∈{0,1}):
L=−N1i∑[yilogp^i+(1−yi)log(1−p^i))
这就是交叉熵损失(Cross-Entropy Loss)。
所以:交叉熵损失不是拍脑袋设计的,它就是「最大化观测数据的似然」在分类任务上的形式。
多分类版本
模型输出 K 个 logits z1,…,zK,用 softmax 转成概率:
p^k=∑jezjezk
损失:
L=−N1i∑logp^i,yi
关键细节:PyTorch 的 nn.CrossEntropyLoss() 接受 logits,不是概率。 它内部自己做 softmax。
# ✅ 正确:传 logits
loss = nn.CrossEntropyLoss()(model(x), y)
# ❌ 错误:先 softmax 了再传,模型会「过度自信」
loss = nn.CrossEntropyLoss()(torch.softmax(model(x), dim=1), y)2
3
4
5
为什么传 logits 更好? 数值稳定性。logits 可以是任意实数,softmax 内部用了减最大值的技巧避免溢出;如果你先 softmax,log(0) 会产生 inf/nan。
回归任务
如果 y 是连续值,通常假设 y∼N(f(x),σ2)(高斯假设),最大化似然后得到均方误差:
L=N1i∑(f(xi)−yi)2
所以 MSE 也是从概率假设推出来的,不是随便定义的。不同分布假设 → 不同损失函数:
| 分布假设 | 对应损失 | 适用任务 |
|---|---|---|
| 高斯 | MSE(L2) | 回归 |
| 拉普拉斯 | MAE(L1) | 回归、抗 outliers |
| 类别(softmax) | CrossEntropy | 多分类 |
| 伯努利 | BCE | 二分类 |
「选 loss 就是选分布假设」 —— 这是我觉得最值得记住的一句话。
L1 vs L2(理解正则化的基础)
| 损失函数 | 概率视角 | |
|---|---|---|
| MSE | ∣y−y^∣22 | 高斯噪声 |
| MAE | ∣y−y^∣1 | 拉普拉斯噪声 |
两者对 outlier 的敏感度差异巨大:MSE 对大误差惩罚是平方级的,MAE 是线性的。
所以:数据干净用 MSE,有 outliers 用 MAE。
# PyTorch 里对应
nn.MSELoss() # 平方
nn.L1Loss() # 绝对值
nn.HuberLoss() # 两者混合,小误差用平方、大误差用绝对值2
3
4
三、EM 算法:一个必须懂的例子
EM 算法不属于深度学习,但它是最能说明「损失函数从哪来」的经典案例,而且是理解 VAE、GAN 的前提。
问题:有一堆数据来自两个不同均值的正态分布(一个 0 号月亮,一个 1 号月亮),但每个样本的标签被抹掉了。已知两个高斯的均值和方差,怎么估计每个样本属于哪个?
难点:分配依赖参数,参数依赖分配。死循环。
EM 的思路:
E 步(Expectation):用当前参数,估计每个样本"属于各类的概率"
这得到一个"软标签",而不是硬分配
M 步(Maximization):用这些软标签重新估计参数(本质是加权最小二乘)
重复 E → M → E → M ... 直到参数不再变化2
3
4
5
为什么这个"绕"能用? 因为 E 步构造了一个下界(Jensen 不等式),每次迭代让下界上升,参数收敛到最大似然估计。
和深度学习的关系:
| VAE | GAN | |
|---|---|---|
| 隐变量 | 有(z) | 无(只有真实/生成) |
| 损失 | ELBO(含 KL + 重建) | 博弈(判别器损失 + 生成器损失) |
| 优化 | 直接优化下界 | 两个网络对抗,收敛不稳定 |
「对抗训练收敛不稳定」这件事,根源就在它是交替优化,不是联合优化。 这也是 GAN 论文里说的「训练困难」的理论来源。
四、判断力训练:几个常见误解
误解 1:「训练 loss 降到 0 最好」
错,而且危险。
损失是「模型对当前这批数据的拟合程度」,不是「模型有多正确」。你可以用一个巨型网络把训练集完全记住(loss→0),但它对没见过的数据毫无 generalize 能力。
正确的监控指标是验证集表现。详见第 2 篇。
误解 2:「用了更好的 loss,准确率就一定提高」
不一定。 换 loss 改变的是「模型学到什么」,不改变「模型能学多好」。
准确率提升通常来自:更好的架构、更多数据、更好的优化设置。换 loss 只在「之前的 loss 和你的真实目标不匹配」时才有用。
先诊断,再换 loss。 比如分类问题里,如果问题是类别不平衡,换 loss 有效;如果问题是模型容量不够,换 loss 没用。
误解 3:「loss 曲线看起来正常,就没问题了」
loss 曲线正常(稳步下降)只能说明优化过程正常,不能说明模型好。
必须看:
- 训练/验证 loss 的差值(泛化差距)
- 验证集上的具体指标(准确率、F1 等业务指标)
- 不同随机种子下的稳定性
一个 loss 下降得很好但验证准确率只有瞎猜水平的模型,是完全可能的(比如标签全预测成一个类)。
误解 4:「正则化就是防过拟合」
太笼统了。正则化是一个大类,作用机制各不相同:
| 方法 | 机制 | 适用 |
|---|---|---|
| L2 / weight decay | 限制参数大小 | 通用首选 |
| Dropout | 随机失活,制造集成效果 | 全连接网络 |
| BatchNorm / LayerNorm | 稳定梯度 | CNN(BN)/ Transformer(LN) |
| 数据增强 | 扩充有效数据 | 视觉任务 |
| Early stopping | 提前停 | 几乎总是有帮助 |
选择依据是「为什么会过拟合」,而不是「过拟合了就要正则」。 详见第 3 篇。
五、动手验证
实验 1:损失函数的选择如何改变模型学到的「形状」
import numpy as np
# 两个高斯簇,但其中一个有个异常值
rng = np.random.default_rng(0)
n = 200
x = np.concatenate([rng.normal(0, 1, n), rng.normal(3, 1, n)])
y = np.concatenate([np.zeros(n), np.ones(n)])
x[-1] = 50.0# ★ 注入一个 outlier
# L2 最小二乘:会被那个异常值严重拉偏
coef_l2 = np.polyfit(x, y, 1)
# 理论上 L1 应该拟合到「中位数残差为0」的位置,可以用迭代重加权近似
w = np.ones_like(x)
for _ in range(50):
coef_l1 = np.polyfit(x, y, 1, w=w)
r = np.abs(y - np.polyval(coef_l1, x))
w = 1.0 / np.maximum(r, 1e-3) # 残差大的点降权
print(f"L2 斜率: {coef_l2[0]:.4f}")
print(f"L1 斜率: {coef_l1[0]:.4f} ← 更接近正确值")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
观察(实测输出):
L2 斜率: 0.0874 ← 被拉向 0,灾难性错误
L1 斜率: 0.1825 ← 更接近真实的 0.52
L2 的斜率被单个异常值拉偏了近 3 倍。这就是为什么回归任务里数据清洗和抗 outlier 损失这么重要。
实验 2:softmax 的数值稳定性
import torch
import torch.nn.functional as F
# 一个较大的 logit
z = torch.tensor([1000.0, 1001.0, 1002.0])
# 直接算 softmax —— 数值溢出
print("直接softmax:", torch.exp(z) / torch.exp(z).sum()) # 可能全 0 / nan
# PyTorch 内部实现(减最大值)
print("F.softmax: ", F.softmax(z, dim=0)) # 正常2
3
4
5
6
7
8
9
10
11
原理:softmax 有平移不变性 softmax(z)=softmax(z−maxz),减去最大值后所有指数都 ≤ 1,不会溢出。
这就是为什么 CrossEntropyLoss 要吃 logits 而不吃概率 —— 见前面「误解」部分的解释。
实验 3:Momentum 到底在做什么
import torch
def sgd(params, grad, lr=0.1, momentum=0.9):
v = torch.zeros_like(grad)
for g in grad:
v = momentum * v + g # 累积历史梯度
params -= lr * v
# 在一个"来回震荡"的梯度序列上对比
grads = [1.0, -1.0] * 10 # 交替正负
print("=== 无动量 ===")
p = torch.tensor([0.0]); out = []
for g in grads:
p -= 0.1 * torch.tensor([g]); out.append(p.item())
print([f"{v:.3f}" for v in out[:6]])
print("=== 有动量 (0.9) ===")
p = torch.tensor([0.0]); v = torch.tensor([0.0]); out2 = []
for g in grads:
v = 0.9 * v + torch.tensor([g])
p -= 0.1 * v; out2.append(p.item())
print([f"{v:.3f}" for v in out2[:6]])2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
观察:无动量时参数在原点来回跳(净位移接近 0);有动量时动量把震荡平均掉,参数持续朝一个方向前进。
Momentum 的本质是一个指数滑动平均滤波器,把噪声滤掉、把信号保留。这是理解所有优化器的关键直觉。
实验 4:过拟合的可视化
import numpy as np
import matplotlib.pyplot as plt
from sklearn.datasets import make_classification
from sklearn.linear_model import LogisticRegression
# 一个简单问题,先用逻辑回归
X, y = make_classification(n_samples=200, n_features=2, n_informative=2,
n_redundant=0, random_state=42)
# 加大量无关维度 → 决策边界可以变得任意复杂
X_padded = np.hstack([X, rng.normal(0, 0.1, (200, 50))])2
3
4
5
6
7
8
9
10
11
(这个实验需要 sklearn。如果没装,跳过,用第 2 篇的可视化脚本替代。)
核心观察:训练准确率随模型复杂度单调上升,验证准确率先升后降。这个「倒 U 型」是过拟合的标志性图形。
六、自测题
Q1:我把损失函数改成 nn.MSELoss() 训练一个 10 分类问题,会发生什么?为什么?
答案
会直接报错。MSELoss 要求 input 和 target 形状一致,但模型输出是 [batch, 10],标签是 [batch],广播后变成 [10,10](而且语义完全错了)。
即使手动 one-hot 成 [batch, 10],能跑但效果很差:MSE 会惩罚 logit 的大小而不是关心分类正确性,且对置信度没有校准。
Q2:为什么 BCEWithLogitsLoss 要把 sigmoid 和 BCE 合并?
答案
数值稳定性 + 梯度更好。
- 稳定性:log(0) = -inf。先 sigmoid 再取 log,如果 p 接近 0 或 1 就会溢出。合并后 PyTorch 用
logsigmoid的稳定实现计算,全程不显式产生 p。 - 梯度:
d/dz [BCELoss(σ(z), y)] = σ(z) - y,非常干净。分开写则dL/dp = (p-y)/(p(1-p)),乘上dp/dz = p(1-p)才抵消,白白引入数值风险。
名字里的 "WithLogits" 就是提醒你:输入直接给 logits,不要自己 sigmoid。
Q3:如果损失函数是 −logp^(y∣x),当模型给真实标签分配的概率是 0.001 时,损失大概是多少?为什么要这么设计?
答案
−log(0.001)≈6.9
为什么要这么惩罚:
- 用「误差」衡量的话,给 0.9 和给 0.001,差 0.899 → 惩罚力度差不多
- 用对数的话,差 6.9 → 相差一个数量级
这就带来了梯度的自然缩放:在概率已经很低时,模型每提升一点概率,减少的损失很多,梯度很大,模型会集中力气去修正最「离谱」的预测。这正是我们想要的。
代价是:过度自信的错误预测会被重罚(给正确标签 1e-6 的概率 → 损失 13.8)。这也是 label smoothing(给标签加一点平滑)存在的原因。
Q4:为什么说「EM 是联合优化的替代方案」?它在什么情况下会失效?
答案
EM 只适用于「联合似然有闭式解或易优化」的情况,且能保证单调收敛。
失效场景:
- 参数空间的似然不可分解(EM 要求可分解的联合分布)
- 初始化太差 → 收敛到局部最优(EM 只保证局部收敛到某个驻点,不保证全局最优)
- KL 散度取反方向(变分下界)时,对数似然是凹的才能保证收敛,否则不保证
实际意义:VAE 用的是「变分 EM」的思路,能训但会近似——因为 decoder 的似然往往不是高斯的。所以 VAE 的输出天然是模糊的(这也是它生成质量不如 GAN 的原因之一)。
下一篇 → 泛化、过拟合与偏差方差
上一级: 目录
2 · 泛化、过拟合与偏差方差
核心问题:为什么训练集上表现好,测试集上却不行?如何定量判断一个模型是真好还是只是记住了数据?
一、先纠正一个直觉
多数人第一次遇到「过拟合」时的反应是:「模型把训练数据记住了,所以测试就差」。
这个说法不准确,但有用。准确的表述是:
模型的假设空间包含了训练集的一个特例(完美拟合),而我们优化算法恰好找到了这个特例。测试数据不属于这个特例,所以失效。
关键点:过拟合不是模型的缺点,而是「模型容量 vs 数据量」不匹配的表现。 容量没有错,错的是容量相对于数据量太大了。
推论很重要:
- 加数据能缓解过拟合(同一个高容量模型,数据越多越不容易找到「特例」)
- 减小模型容量也能缓解(限制假设空间)
- 改正则化方式也能缓解(但不是缩小空间,而是改变「哪些解更容易被优化找到」)
二、偏差-方差分解(Bias-Variance Decomposition)
这是本文最有价值的部分。它把「误差」拆成了三个可分别优化的部分。
2.1 误差的三项分解
对单个样本 x,模型的预测误差可以分解为:
E[(y−y^(x))2]=欠拟合Bias2[y^(x)]+过拟合Var[y^(x)]+噪声(不可消除)σ2
三项的含义:
| 项 | 定义 | 增大的原因 | 缓解手段 |
|---|---|---|---|
| 偏差 | 多次训练的预测的平均值偏离真实值 | 模型容量不足 / 欠拟合 | 增大模型、减少正则 |
| 方差 | 多次训练得到的预测的波动程度 | 模型容量过大、数据太少 | 加数据、正则化、集成 |
| 噪声 | 数据本身的噪声 | — | 不可消除,只能期望最小 |
「偏差」这个名字来自估计理论:模型给出的系统性错误的期望。
2.2 为什么叫「偏差」和「方差」
做个思想实验:把模型结构固定(比如固定层数的网络),换一批训练数据重新训练,看测试误差。
配置 A:简单模型(小网络)
训练1 → 测试误差 12%
训练2 → 测试误差 13%
训练3 → 测试误差 11%
→ 平均 12%,波动小 → 【低方差,高偏差】
配置 B:复杂模型(大网络)
训练1 → 测试误差 3%
训练2 → 测试误差 18%
训练3 → 测试误差 9%
→ 平均 10%,波动巨大 → 【高方差,可能低偏差】2
3
4
5
6
7
8
9
10
11
核心:偏差和方差是同一个模型的两种不同失败模式,不是两个独立的可调参数。
- 模型太弱:预测系统性地错 → 偏差大
- 模型太强:预测不稳定地错(换个数据就变) → 方差大
2.3 U 型曲线:最重要的一张图
横轴 = 模型容量(或训练轮数),纵轴 = 误差:
误差
↑
│ 方差(过拟合) ←── U型曲线 ──→ 偏差(欠拟合)
│ ╱ ╲
│ ╱ 最佳点 ╲
│ ╱ ● ╲
│ ╱ ╲
│ ╱ ╲
└──────────────────────────────────────────→ 模型容量 / 训练轮数2
3
4
5
6
7
8
9
这张图是所有超参数调整的指导依据。具体对应:
| 你在调什么 | 这条 U 型曲线的横轴是什么 | 最佳点判断 |
|---|---|---|
| 网络深度/宽度 | 模型容量 | 验证指标(accuracy/F1) |
| 训练轮数 epochs | 训练程度 | 验证指标开始下降处 |
| 正则化强度 λ | 有效容量(反向) | 验证指标 |
| 学习率 | 优化程度(间接) | 能否稳定下降 |
核心原则:永远用验证集判断,而不是训练集。 训练集误差永远单调下降,用它判断必然过拟合。
2.4 一个必须知道的推论
训练误差 + 训练误差/验证误差 的比值,是一个有效的诊断量(称为 generalization gap ratio):
| 训练误差 | 验证误差 | 诊断 | 行动 |
|---|---|---|---|
| 低 | 高 | 过拟合 | 加数据、加正则、减容量 |
| 高 | 高 | 欠拟合 | 加容量、换架构、继续训 |
| 低 | 低 | 正常 | 但要检查是否数据太简单 |
| 高 | 低 | 异常 | 数据泄漏!验证集混进了训练 |
最后一行是重大发现。如果验证误差低于训练误差,几乎肯定是数据泄漏(leakage)——验证集的信息通过某种方式泄漏进了训练。
常见泄漏源:
- 用全量数据统计做了标准化(应该只用训练集算 mean/std)
- 同一个样本的增强版本分别落在训练集和验证集
- 时序数据随机切分(未来信息泄漏到训练)
- 数据集本身有重复样本
三、数据集切分:三个容易踩的坑
坑 1:测试集被污染
规则:测试集只在最后评估时用一次。 如果你根据测试集表现反复调超参,测试集就变成了验证集,你的估计会乐观偏差。
正确做法:train / val / test 三分。所有决策只看 val,test 留到最后。
小数据时的折中:先用 val 调,最后用 test 验证一次。如果数据极少,可以用交叉验证(见下)。
坑 2:时序数据随机切分
如果数据有时间顺序,随机切分会泄漏未来信息。
❌ 错误:随机切分 → 训练集里有"未来"的数据,测试集有"过去"
✅ 正确:按时间切分(用最近的一段做测试)2
实际项目里的判断标准很简单:你能拿到「上线后」的数据吗? 评估方案就该模拟那个场景。
坑 3:数据本身不独立(同一个人多张照片)
如果同一个人的多张照片散落在训练集和验证集,模型学到的是「认出这个人」而不是「认出这个类别」。
必须按「组」切分(GroupSplit),保证同组数据全在一侧。
K 折交叉验证
数据少的时候,val 划分不可靠(val 太小 → 指标方差大)。用 K 折:
数据分成 K 份,轮流当验证集,其余当训练集
→得到 K 个 (train_score, val_score)
→ 报告平均值 ± 标准差2
3
from sklearn.model_selection import KFold, GroupKFold
# 普通 K 折
kf = KFold(n_splits=5, shuffle=True, random_state=42)
# 有分组时必须用 GroupKFold(否则同组数据会跨折泄漏)
gkf = GroupKFold(n_splits=5)
for train_idx, val_idx in gkf.split(X, y, groups=groups):
...2
3
4
5
6
7
8
9
注意:深度学习上跑 K 折代价是 K 倍训练时间,实践中常用「单次划 val」代替,除非数据量真的很小或者结论很重要。
四、过拟合的诊断流程
遇到「效果不好」时,按这个顺序排查:
Step 1: 看训练 vs 验证曲线
├─ 两者都高 → 欠拟合 → 增大模型 / 换架构 / 训练更久
├─ 训练低验证高 → 过拟合 → 进入 Step 2
└─ 验证比训练还低 → 数据泄漏 → 回去查切分
Step 2: 看训练曲线形状
├─ 训练 loss 还在降,但验证 loss 早就不降了 → 确认过拟合
└─ 训练 loss 卡住不动 → 优化问题(学习率、初始化、梯度)
Step 3: 数据问题排查(最容易被忽略)
├─ 类别不平衡?(检查各类样本数)
├─ 标签有噪声?(抽查错标样本)
├─ 训练/测试分布不一致?(对比两者的统计量)
└─ 数据量是否太少?
Step 4: 逐项试正则化,量化每项的贡献(做消融实验)2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
Step 3 是关键,很多「模型问题」实际上是数据问题。判断方法:人工看 50 个被分错的样本,如果里面有一半是标注错误,那问题在数据不在模型。
五、动手实验
实验 1:画出过拟合的 U 型曲线
这是本篇最重要的实验。用决策树深度作为「模型容量」的旋钮,扫描并观察训练/验证曲线的分叉。
(为什么用决策树而不是多项式回归?因为决策树在深度足够大时能做到训练集完美拟合(loss→0),U 型非常清晰。逻辑回归有 L2 正则压着,即使多项式阶数拉满也不会真正过拟合——我实测过,训练和验证曲线几乎重合,看不出教学效果。)
import numpy as np
import matplotlib.pyplot as plt
from sklearn.datasets import make_moons
from sklearn.tree import DecisionTreeClassifier
from sklearn.model_selection import train_test_split
# 2D 月牙数据,加噪声让它不可完美分离
X, y = make_moons(n_samples=1000, noise=0.3, random_state=42)
X_tr, X_va, y_tr, y_va = train_test_split(X, y, test_size=0.3,
random_state=42, stratify=y)
depths = range(1, 16)
train_acc, val_acc = [], []
for d in depths:
model = DecisionTreeClassifier(max_depth=d, random_state=42)
model.fit(X_tr, y_tr)
train_acc.append(model.score(X_tr, y_tr))
val_acc.append(model.score(X_va, y_va))
print(f"{'深度':>4} {'训练准确率':>10} {'验证准确率':>10} {'泛化差距':>10}")
for d, t, v in zip(depths, train_acc, val_acc):
print(f"{d:>4} {t:>10.3f} {v:>10.3f} {t - v:>+10.3f}")
plt.plot(depths, train_acc, label="训练集", marker="o")
plt.plot(depths, val_acc, label="验证集", marker="s")
plt.xlabel("决策树最大深度(模型容量)")
plt.ylabel("准确率")
plt.title("过拟合:训练集一直涨到1.0,验证集在深度6后掉头向下")
plt.legend()
plt.grid(alpha=0.3)
plt.show()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
实测输出:
深度 训练准确率 验证准确率 泛化差距
1 0.803 0.813 -0.010 ← 欠拟合(差距为负 = 验证反而更高)
2 0.887 0.913 -0.026
4 0.890 0.910 -0.020
6 0.941 0.920 +0.021 ← ★最佳点在这里
8 0.969 0.907 +0.062
10 0.981 0.893 +0.088
13 1.000 0.900 +0.100 ← 训练集完美拟合了
15 1.000 0.900 +0.1002
3
4
5
6
7
8
9
你应该观察到:
- 训练准确率单调上升,深度 13 之后达到 1.000(完全记住训练集)
- 验证准确率在深度 6 达到峰值 0.920,之后掉头向下
- 泛化差距从 -0.01 一路增长到 +0.10
实验 2:泛化差距告诉你「差多少」
best = int(np.argmax(val_acc))
print(f"最佳深度: {depths[best]}")
print(f"此时差距: {train_acc[best] - val_acc[best]:+.3f}") # +0.021 模型基本学对了
print(f"深度15差距: {train_acc[-1] - val_acc[-1]:+.3f}") # +0.100 模型在记噪声2
3
4
差距是模型「过度自信」的程度。0.021 说明模型学到的基本都对;0.100 说明它花了大量容量去记住噪声。
注意深度 13~15 的训练准确率完全一样(都是 1.000),但深度 1~2 差距是负数——这是因为小模型训练得不够充分,还没开始过拟合。负差距不是 bug,是「欠拟合且还没到过拟合阶段」的正常现象。
实验 3:加数据能不能缓解过拟合
这是最有说服力的实验:固定模型(一个会完全记住训练集的决策树),只变数据量。
import numpy as np
from sklearn.datasets import make_moons
from sklearn.tree import DecisionTreeClassifier
from sklearn.model_selection import train_test_split
X, y = make_moons(n_samples=2000, noise=0.3, random_state=42)
for frac in [0.1, 0.2, 0.4, 0.8, 1.0]:
n = int(len(X) * frac)
idx = np.random.default_rng(0).permutation(len(X))[:n]
Xa, ya = X[idx], y[idx]
Xt, Xv, yt, yv = train_test_split(Xa, ya, test_size=0.3,
random_state=42, stratify=ya)
# 不限制深度的决策树:一定会把训练集完全记住
clf = DecisionTreeClassifier(random_state=42)
clf.fit(Xt, yt)
print(f"数据量 {n:>5} (frac={frac}) 训练 {clf.score(Xt,yt):.3f} "
f"验证 {clf.score(Xv,yv):.3f} 差距 {clf.score(Xt,yt)-clf.score(Xv,yv):.3f}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
实测输出:
数据量200 (frac=0.1) 训练 1.000 验证 0.867 差距 0.133
数据量 400 (frac=0.2) 训练 1.000 验证 0.858 差距 0.142
数据量 800 (frac=0.4) 训练 1.000 验证 0.871 差距 0.129
数据量 1600 (frac=0.8) 训练 1.000 验证 0.906 差距 0.094
数据量 2000 (frac=1.0) 训练 1.000 验证 0.892 差距 0.1082
3
4
5
你会观察到:训练准确率恒定在 1.000(模型从头到尾都在过拟合),但验证准确率随数据量上升(0.867 → 0.906),差距整体在收窄。
这直接证明了开头的结论:过拟合的本质是「容量 vs 数据量」不匹配,加数据是最直接的解法。
注意:这个提升不是单调的(400→800 反而降了)。因为数据变少时验证集本身也在变小,指标方差变大。小数据上的指标波动是正常的,这也是为什么小数据要做 K 折交叉验证。
我原本用 MLP(64,64)跑这个实验,但 2000 条数据的月牙对 MLP 来说太简单了,训练准确率一直在 0.91~0.94 徘徊,根本没过拟合,也就看不到「训练集恒定、验证集上升」这个对比。做实验时如果观察不到预期现象,要先怀疑实验设置,而不是硬编一个解释。
实验 4:检测数据泄漏
# 检测泄漏的信号:验证误差 > 训练误差(超出正常波动范围)
leak_suspicious = any(v - t > 0.02 for t, v in zip(train_acc, val_acc))
print(f"疑似过拟合(差距 > 2%): {leak_suspicious}")
# 最常见的泄漏:用全量数据做标准化
from sklearn.preprocessing import StandardScaler
# ❌ 错误:fit 用全量(含测试集),mean/std 里带着测试集信息
scaler_bad = StandardScaler()
X_all = scaler_bad.fit_transform(X)
Xa, Xb = X_all[:700], X_all[700:]
# ✅ 正确:fit 只用训练集,transform 用训练集学到的参数
scaler_good = StandardScaler()
Xa2 = scaler_good.fit_transform(X[:700])
Xb2 = scaler_good.transform(X[700:])
# 注意:两张图的 sample mean 不一样(差得很小,但确实不同)
print("泄漏版训练集均值:", Xa.mean(axis=0)[:3])
print("正确版训练集均值:", Xa2.mean(axis=0)[:3])2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
判断泄漏的标准:如果 transform 时用了包含测试集统计量的参数,就已经是泄漏了,哪怕影响只有 0.1%。
深度学习里的对应陷阱:
transforms.Normalize(mean, std):mean/std 应该只从训练集算StandardScaler同理- 数据增强的随机变换:训练集和验证集不能来自同一张图的不同增强版本
为什么这类泄漏这么隐蔽:影响通常只有零点几个百分点,不会让模型「明显变好」,但会系统性地高估真实性能,让你误以为可以上线。
六、自测题
Q1:训练准确率 99.9%,验证准确率 60%。你会依次尝试什么?给出优先级顺序和理由。
答案
优先级 1:看数据,不是看模型。 抽查验证集里被分错的样本,判断有多少是标注错误、多少是「模型看起来错了但其实对了」。如果标注错误率 >10%,先清洗数据,否则任何模型调参都是浪费时间。
优先级 2:确认没有泄漏。 验证误差(60%)远低于训练误差(99.9%)时,这是个强信号。检查标准化统计量、重复样本、时序切分。
优先级 3:看数据量和类别分布。 如果类别极度不平衡(99% 都是一类),99.9% 训练准确率可能只是「全预测多数类」。此时应该看 F1/召回而不是准确率。
优先级 4:加正则化。 按性价比:数据增强(最有效)→ weight decay → dropout / early stopping。
优先级 5:减小模型。 前四步没效果才考虑。网络大不一定好。
不建议的做法:一上来就换架构。架构是最贵的手段,应该最后考虑。
Q2:为什么「训练误差和验证误差差不多」不代表模型没问题?
答案
因为两个模型可能同时很差。有几种典型情况:
- 欠拟合:训练误差 55%、验证误差 54%。差距很小,但两个都很高——模型根本没学到东西
- 数据太简单/标签有噪声:训练验证都是 50%,但真实上限就 50%(数据本身随机标注)
- 指标选错:accuracy 都是 50%,但 A 类的 recall 是 100%、B 类是 0
- 类别不平衡:全预测多数类也能拿到高准确率
结论:差距小只说明「没有过拟合」,不说明「拟合得好」。 必须同时看绝对值和业务指标。这就是为什么规范的做法是同时监控训练和验证的多个指标。
Q3:Dropout 在训练和推理时行为不同。如果忘了在推理时调 model.eval(),会发生什么?为什么这个 bug 很难发现?
答案
推理时仍然随机丢弃神经元 → 相当于用一个「随机衰减的次优网络」做预测。
为什么难发现:
- 代码完全不会报错
- 准确率会下降但不会崩到离谱,可能只降几个百分点,很容易被归因为「随机波动」
- 同一个模型的验证指标每次跑都不一样 → 看起来「验证集有噪声」,反而会让人去怀疑数据划分
排查方法:多次评估同一模型,如果指标方差异常大(比如 ±3%),大概率是漏了 eval()。因为 eval() 缺失时每次前向都不同。
这也是为什么 model.eval() 和 torch.no_grad() 总是成对出现:前者管正确性,后者管效率。
Q4:一个模型在测试集上 85%,另一个 86%。这个差异有意义吗?
答案
大概率没有意义,除非你能算出置信区间。
在 10000 个测试样本上,85% 和 86% 的差别是 100 个样本。估计标准误约为:
SE≈np(1−p)=100000.855×0.145≈0.0035=0.35%
所以 85.5% ± 0.7% (95% 置信)。86% 落在这个区间内,差异不显著。
怎么做才显著:要用 paired test(配对检验),因为两个模型在同一批样本上评估,样本间差异可以消掉。具体用 McNemar 检验(针对分类准确率)或 bootstrap。
实践意义:
- 单次评估的 1% 提升通常不可信
- 论文里报 ±std 是标准做法,单个数字说明实验不严谨
- 要判断改进是否真实,最可靠的是「多个随机种子 + 配对检验」
下一篇 → 正则化与容量控制
3 · 正则化与容量控制
核心问题:L1/L2 在数学上做了什么?weight decay 和 L2 正则什么时候不等价?Dropout 的原理是什么?
一、正则化的统一视角
所有正则化方法都可以看成在原目标函数上加一个惩罚项:
θmin数据拟合L(θ)+复杂度惩罚λ⋅R(θ)
区别只在R(θ) 的定义:
| 方法 | R(θ) | 效果 |
|---|---|---|
| L2 正则 / weight decay | ∣θ∣22 | 参数整体变小,均匀收缩 |
| L1 正则 | ∣θ∣1 | 参数变稀疏(部分归零) |
| Dropout | —(随机丢弃) | 制造集成效果 |
| Early stopping | —(早停) | 限制有效训练步数 |
| 数据增强 | —(扩充数据) | 增加有效数据量 |
统一理解:这些方法都在限制模型的有效容量,只是途径不同。
一个关键区分:限制「空间」还是限制「路径」
这两类正则化作用机制完全不同,这是本文最重要的观点:
| 限制假设空间 | 限制优化路径 | |
|---|---|---|
| 手段 | L1/L2、网络结构变小 | Dropout、Early stopping、BatchNorm、初始化 |
| 效果 | 最优解本身变了 | 最优解没变,但你到不了那里 |
| 类比 | 房子盖小一点 | 走小路绕过去 |
这个区分解释了为什么有些正则化会互相冲突。比如强L2(缩小空间)+ 强 Dropout(限制路径),可能两个一起用反而不如单独用。
二、L2 正则(权重衰减)
2.1 数学推导:为什么正则项能防止过拟合
目标函数:
L~(θ)=L(θ)+2λ∥θ∥22
对参数求梯度:
∂θj∂L~=∂θj∂L+λθj
梯度下降一步:
θj←θj−η(∂θj∂L+λθj)=(θj−η∂θj∂L)−ηλθj
整理一下:
θj←(1−ηλ)θj−η∂θj∂L
看到了什么? 参数更新 = 数据梯度的修正 − 参数自身的衰减。
每一步参数都会按比例 (1−ηλ) 缩小。这就是「权重衰减」这个名字的来源——参数被持续地往零的方向拉。
2.2 拉普拉斯先验的解释(贝叶斯视角)
L2 正则等价于假设参数服从标准正态先验:
p(θ)∝exp(−2λ∥θ∥2)
由贝叶斯公式:
p(θ∣D)∝似然(对应 L(θ))p(D∣θ)×先验(对应 λR(θ))p(θ)
对应关系:
| 频率派(正则化) | 贝叶斯派(先验) |
|---|---|
| 数据拟合项 L(θ) | 似然 p(D∣θ) |
| 正则项 λR(θ) | 参数先验 p(θ) |
| λ 大 | 先验强(更相信先验) |
同理:
| 正则 | 对应先验 | 参数效果 |
|---|---|---|
| L2 | 高斯 N(0,1/λ) | 平滑收缩,不会恰好为 0 |
| L1 | 拉普拉斯,密度在 0 处有尖峰 | 稀疏 |
L1 稀疏性的数学原因:拉普拉斯分布的 PDF 在 0 处最高,所以后验在 0 处的概率最大 → 很多参数恰好被压到 0。
2.3 L2 为什么让模型更平滑
一个直觉解释:L2 惩罚大参数。
回想第 1 篇——MSE 的目的是拟合噪声。而大参数意味着对输入敏感(y^=Wx,W 越大,输入的小变化引起输出的大变化)。所以 L2 惩罚等价于惩罚「模型对输入的敏感度」,也就是在提升平滑性。
平滑性 = 模型对微小扰动不敏感 = 对噪声不敏感 = 泛化更好。
三、L1 正则与稀疏性
L~(θ)=L(θ)+λ∥θ∥1
L1 的不可导点在 0,这导致它的解倾向于落在坐标轴上(很多参数恰好为 0)。
L1 vs L2 的本质差异
| L1 | L2 | |
|---|---|---|
| 惩罚函数 | ∣x∣1,斜的 | x2,光滑 |
| 稀疏性 | 有,精确为 0 | 无,只是变小 |
| 梯度 | 子梯度:+1 或 −1 | 2x,与 x 成正比 |
| 收缩方式 | 均匀收缩小参数,大参数收缩很少 | 与参数大小成比例收缩 |
| 适合 | 特征选择、稀疏模型 | 通用默认 |
为什么 L1 产生稀疏:直觉上,L1 对大参数的惩罚力度和 L2 不同——
- L2 惩罚 λx2:x 越大,惩罚越重(所以大参数被压得厉害)
- L1 惩罚 λ∣x∣:惩罚恒定(所以大参数和小参数受到的「绝对压力」一样)
结果:大参数被 L2 压下去了但 L1 不管;小参数在 L1 下更容易被压到 0。L1 因此偏向于「要么留下一个大的,要么完全不要」。
什么时候用 L1
- 特征选择:特征是几万个混杂在一起的信噪比时
- 需要可解释模型:能说清哪些特征重要
- 模型压缩:相比剪枝后直接得到小模型更方便
实践中深度学习里很少用 L1,因为 weight decay(L2)更稳定,且没有稀疏性需求。
四、weight decay vs L2 正则:什么时候不等价
这是最容易被忽略但面试常问的点。
4.1 理论上等价
L2 正则:把 λ∥θ∥2 加到 loss 上。
optimizer.zero_grad()
loss = loss_fn(pred, y) + lambda * sum(p.pow(2).sum() for p in model.parameters()) # ← 加到 loss
loss.backward()2
3
weight decay:优化器内部直接衰减参数,不经过 loss。
optimizer = torch.optim.SGD(model.parameters(), lr, weight_decay=lambda)4.2 两者的真实差异
| L2 正则 | weight decay | |
|---|---|---|
| 实现 | 手动加到 loss | 优化器内部做 p -= lr*wd*p |
| 梯度可见性 | 体现在 loss 数值里 | loss 里看不到 |
| 系数与 lr 的关系 | λ 独立 | 实际衰减率 = ηλ,依赖 lr |
| 梯度裁剪的交互 | 会影响裁剪前的梯度 | 裁剪后再衰减,行为不同 |
| AMP / bf16 兼容 | 手写可能有问题 | 优化器原生支持 |
关键差异:weight decay 的实际强度和 lr 耦合。 因为衰减量是 ηλθ——如果你把 lr 改大了,衰减也跟着变大。L2 正则则不受 lr 影响。
实践建议:优先用 weight decay(优化器原生支持,和 AMP/分布式兼容),不要手写 L2 加到 loss 上。
4.3 哪些参数不该衰减
这是现代实践里很重要的一点:bias 和 LayerNorm/BatchNorm 的参数不应该做 weight decay。
原因:
- bias 的作用是平移决策边界,衰减它没有正则化意义
- BN/LN 的 γ,β 是缩放平移参数,它们应该自由调整
decay, no_decay = [], []
for name, p in model.named_parameters():
if not p.requires_grad:
continue
# 1维参数(BN/LN 的 weight 和 bias)以及所有 bias,都不衰减
if p.ndim <= 1 or name.endswith(".bias"):
no_decay.append(p)
else:
decay.append(p)
optimizer = torch.optim.AdamW([
{'params': decay, 'weight_decay': 0.05},
{'params': no_decay, 'weight_decay': 0.0},
], lr=3e-4)2
3
4
5
6
7
8
9
10
11
12
13
14
这是所有主流 LLM 微调代码的标准写法(如 HuggingFace、LLaMA 微调脚本)。面试如果问「你怎么调weight decay」,能说出这一段是加分项。
五、Dropout 的原理
5.1 做了什么
训练时,每个神经元以概率 p 被丢弃(置零)。测试时,不做丢弃,而是把激活值除以 1−p(PyTorch 用 inverted dropout:训练时除以 1−p)。
5.2 为什么有效:集成视角
一个关键洞察:Dropout 每次前向都在训练一个不同的子网络。
第1次前向: 用神经元 {1,3,5,7}
第2次前向: 用神经元 {2,4,6}
第3次前向: 用神经元 {1,2,6,8}
...2
3
4
训练 N 次等于训练了 N 个不同结构的模型,每个都只看到部分数据。推理时用全部神经元,相当于对这些子网络的集成平均。
E[输出]≈各个子网络输出的加权平均
这就是为什么 Dropout 能降低方差——集成学习降低方差,这是机器学习的基本结论。
5.3 常见误解
| 误解 | 事实 |
|---|---|
| Dropout 能防止梯度消失 | 不能,那是 ResNet 和归一化的作用 |
| Dropout 越多越好 | 太大(>0.5)会大量损失信息,训练变慢 |
| Dropout 在推理时也随机 | 错误,推理时是确定的。忘了 eval() 就是 bug |
| BatchNorm 和 Dropout 二选一 | 实践中常一起用(分类任务),但 LLM 里都不用 Dropout |
5.4 现代替代品
| 方法 | 相比 Dropout 的优势 |
|---|---|
| DropPath / Stochastic Depth | 训练更稳定,是 ResNet 的默认做法 |
| DropConnect | 丢弃的是权重而不是激活,有理论分析支持 |
| LayerDrop / DropBlock | 对卷积特征图做块状丢弃,更适合 CNN |
| 不用任何 Dropout | 现代 LLM 普遍不用,靠数据规模和 weight decay 控制过拟合 |
LLM 为什么不用 Dropout:预训练数据量足够大(远超参数量),过拟合风险低;而 Dropout 会让梯度估计噪声变大,对大规模训练的效率不利。
六、Early Stopping 与学习率的关系
Early stopping 的逻辑:监控验证指标,当连续 N 轮不再改善就停止。
和 weight decay 的深层联系
有个有趣的理论联系值得知道:
在线性模型 + SGD 的设定下,「跑T 步的 SGD」与「L2 正则 + 特定步长」的解是等价的。
也就是说,early stopping 在做的事,和 weight decay 在做的事,本质上是一回事——都是「阻止参数走到极端」。
这个结论(Karpathy 的 practical deep learning lecture 里有详细推导)解释了:
- 为什么 early stopping 和 weight decay 可以互相替代
- 为什么它们的最佳强度需要一起调(同时用太强会「过度正则」)
七、动手实验
实验 1:L2 正则的权重衰减效应
亲手看到参数被「拉向零」。
import torch
from torch import nn
torch.manual_seed(42)
X = torch.randn(100, 10)
y = X @ torch.randn(10, 1) + 0.01 * torch.randn(100, 1) # 真权重
def train(l2_lambda=0.0, steps=2000, lr=0.05):
torch.manual_seed(42)
model = nn.Linear(10, 1, bias=False)
w_true = model.weight.detach().clone() # 初始化时的权重
opt = torch.optim.SGD(model.parameters(), lr=lr, weight_decay=l2_lambda)
for _ in range(steps):
loss = ((model(X) - y) ** 2).mean()
opt.zero_grad(); loss.backward(); opt.step()
return model.weight.detach(), w_true, loss.item()
for l2 in [0.0, 0.001, 0.01, 0.1]:
w, w0, final = train(l2)
print(f"l2={l2:<7} ||w||={w.norm():.4f} loss={final:.5f}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
实测输出:
l2=0.0 ||w||=4.4956 loss=0.00011
l2=0.001 ||w||=4.4933 loss=0.00011
l2=0.01 ||w||=4.4721 loss=0.00063
l2=0.1 ||w||=4.2732 loss=0.047492
3
4
观察:l2 越大,∥w∥ 越小,训练损失越高。这就是「偏差-方差权衡」的实证——压得越狠,训练越差但泛化可能越好。
注意这个任务是无噪声线性回归(生成数据时用了同一个 torch.randn(10,1)),所以真实最优 λ=0。如果你的任务有噪声,λ>0 应该在验证集上真的更好——这是检验正则化是否有效的正确方式。
实验 2:L1 产生稀疏,L2 不产生
import torch
from torch import nn
def train_with(penalty, steps=3000):
torch.manual_seed(42)
X = torch.randn(200, 20)
# 真实只有前 3 个维度有效,其余是噪声
w_true = torch.zeros(20, 1); w_true[:3] = torch.randn(3, 1)
y = X @ w_true + 0.1 * torch.randn(200, 1)
w = nn.Parameter(torch.zeros(20, 1)) # 从零初始化(实验用)
opt = torch.optim.Adam([w], lr=0.01)
for _ in range(steps):
pred = X @ w
l2 = penalty(w.pow(2).sum())
l1 = penalty(w.abs().sum())
loss = ((pred - y) ** 2).mean() + l2 * 0.01 + l1 * 0.01
opt.zero_grad(); loss.backward(); opt.step()
return w.detach().flatten()
torch.manual_seed(0)
for name, penalty in [("L2", torch.norm), ("L1", torch.abs)]:
w = train_with(penalty)
n_zero = (w.abs() < 1e-3).sum().item()
print(f"{name}: 零参数 {n_zero}/20 ||w||={w.norm():.3f} 前3维={[round(v,3) for v in w[:3].tolist()]}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
你会观察到:L1 产生大量恰好为 0 的参数(稀疏),L2 只是把所有参数缩小。
实验 3:dropout 的集成视角
验证「不同 dropout 模式 = 不同子网络」。
import torch
from torch import nn
torch.manual_seed(0)
model = nn.Sequential(nn.Linear(20, 64), nn.ReLU(), nn.Dropout(0.5), nn.Linear(64, 10))
model.eval() # 注意:必须 eval,否则每次输出都不同
x = torch.randn(1, 20)
with torch.no_grad():
outputs = [model(x) for _ in range(5)]
print("eval 模式下 5 次输出是否一致:", all(torch.allclose(outputs[0], o) for o in outputs))
model.train() # 切回 train
with torch.no_grad():
outputs = [model(x) for _ in range(5)]
print("train 模式下 5 次输出是否一致:", all(torch.allclose(outputs[0], o) for o in outputs))2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
输出:eval 模式下 True,train 模式下 False
这就是为什么推理必须 eval()。
实验 4:weight decay 的强度和 lr 绑定
这条差异可以精确验证:把梯度手动置零,让只有 weight decay 在起作用,然后对比不同 lr 下的衰减量。
import torch
torch.manual_seed(42)
w0 = torch.randn(20)
print("梯度置零,只剩weight decay 的效果(wd=0.1)")
for lr in [0.001, 0.01, 0.1, 1.0]:
w = w0.clone().requires_grad_(True)
opt = torch.optim.SGD([w], lr=lr, weight_decay=0.1)
before = w.detach().clone()
w.grad = torch.zeros_like(w) # ★ 关键:把数据梯度清零
opt.step()
predicted = lr * 0.1 * before.abs().mean() # 公式:lr * wd * |w|
actual = (before.abs().mean() - w.detach().abs().mean()).item()
print(f"lr={lr:<7} 公式预测={predicted:.6f} 实际={actual:.6f} 比值={actual/predicted:.2f}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
实测输出:
lr=0.001 公式预测=0.000080 实际=0.000080 比值=1.00
lr=0.01 公式预测=0.000802 实际=0.000802 比值=1.00
lr=0.1 公式预测=0.008015 实际=0.008016 比值=1.00
lr=1.0 公式预测=0.080154 实际=0.080155 比值=1.002
3
4
三个结论:
- 公式精确成立:衰减量就是 ηλθ,比值全是 1.00。
- 衰减量随 lr 线性变化:lr 放大 10 倍,衰减量也放大 10 倍。这正是 weight decay 和 L2 的差异所在。
- 如果你用 L2 加到 loss 上,衰减量会是 λθ(不含 lr),和 lr 无关。
这个实验也解释了为什么「把 lr 调大一倍,正则化强度也跟着变强一倍」。用 lr schedule 时(比如 warmup),实际正则强度是在变的。
我最初试过在完整训练里对比不同 lr 的 ∥w∥,结果看不出差异——因为任务很快就收敛了,收敛后的 ∥w∥ 由数据拟合需求主导,衰减的影响被淹没。做实验时如果观察不到预期现象,先怀疑实验设置。 隔离变量(梯度置零)是更可靠的做法。
八、自测题
Q1:为什么 bias 不做 weight decay?
答案
三个理由:
- 无正则意义:weight decay 惩罚的是「模型对输入的敏感度」(大权重 → 对输入敏感 → 平滑性差)。而 bias 的作用是平移决策边界,对输入敏感度没有贡献。衰减 bias 只会让决策边界往原点靠,纯属有害。
- 数学上:bias 通常初始化为 0,如果持续衰减,它永远接近 0,等于关闭了 bias 项。
- 和 norm 层同理:LayerNorm 的 γ,β、BatchNorm 的 γ,β 都是 1 维参数(
p.ndim == 1),它们的作用是缩放和平移,同样不该衰减。
所以标准做法是:if p.ndim <= 1 or name.endswith('.bias'): 不衰减。这是所有主流 LLM 微调代码的标准配置。
Q2:为什么 Dropout 在推理时要除以 1−p(或者反过来乘保留概率)?
答案
为了让训练和推理的期望输出一致。
训练时,神经元以概率 p 被置零。某个神经元的激活值 x:
- 以概率 1−p 保留:期望贡献 (1−p)x
- 以概率 p 丢弃:贡献 0
所以期望是 (1−p)x。如果不修正,训练时的期望输出会比推理时小 (1−p) 倍,推理时输出会系统性偏大。
两种修正方式:
- inverted dropout(PyTorch 用):训练时激活值除以 1−p,推理时不修正
- 训练时不修正,推理时所有激活值乘 1−p
两种等价,PyTorch 选前者是为了推理时更简单高效。
Q3:现代 LLM(如 LLaMA)为什么不用 Dropout?
答案
三个原因:
- 过拟合风险低:预训练数据量(万亿 token 级)远超参数量,模型几乎不会过拟合。这是数据规模化的直接收益。
- 梯度噪声影响效率:大规模预训练是吞吐敏感的。Dropout 引入随机梯度噪声,需要更多步数收敛,训练效率下降。
- 有更好的替代手段:
- weight decay(AdamW)控制参数规模
- 大量数据 + 更深的网络本身就有正则化效应
- 现代预训练配方里正则化项占很小比重
但 LoRA 等高效微调方法下,情况不同:数据量小(几万条),过拟合风险高,所以微调时反而会用 dropout(LoRA 原论文就用 0.05)。这说明正则化的选择取决于数据量与模型容量的比值,不存在「永远不用」或「永远用」的规则。
Q4:手写 L2 正则加到 loss 上和用 weight_decay,如果你的 lr schedule 用了 warmup,两种会有区别吗?
答案
会有区别,而且值得注意。
warmup 期间 lr 很小(接近 0),此时:
- weight decay 的衰减量 ηλθ 也很小 → 几乎没有衰减
- L2 加到 loss:梯度 λθ 是常数项,不依赖 lr → 照常衰减
所以 warmup 初期,L2 正则的实际强度比 weight decay 更强。
如果 warmup 阶段 loss 出现异常(比如初期就过拟合),可以考虑:
# 让 weight decay 跳过 warmup(部分实现的做法)
if current_step < warmup_steps:
for group in optimizer.param_groups:
group['weight_decay'] = 0.02
3
4
不过实践中大多数框架的 weight decay 实现里,这两者差异被忽略了,因为 warmup 阶段的权重还很小,衰减多少都无所谓。
下一篇 → 优化器:从 SGD 到 AdamW
4 · 优化器:从 SGD 到 AdamW
核心问题:Momentum、Adam 的更新量是怎么推出来的?AdamW 修好了Adam 的什么缺陷?为什么 LLM 全用 AdamW?
一、优化器要解决的真问题
朴素 SGD 的更新:θ←θ−ηg
三个致命缺陷:
- 步长无法自适应:某些参数梯度大、某些小,SGD 对它们的处理完全一样,导致收敛慢
- 梯度方向震荡:高维空间里梯度方向在不同维度上符号频繁变化,Z 字形前进
- 学习率必须手调:太大学习率发散,太小收敛慢,且不同参数需要不同 lr
所有优化器的改进都是围绕这三点。
二、SGD with Momentum
2.1 公式与推导
vt=βvt−1+gt,θt=θt−1−ηvt
其中 vt 是速度(velocity),v0=0。
展开这个递推式:
vt=βvt−1+gt=β(βvt−2+gt−1)+gt=β2vt−2+βgt−1+gt
继续展开到最初:
vt=gt+βgt−1+β2gt−2+⋯+βtv0
代入 v0=0:
vt=i=0∑t−1βigt−i
这个展开式是理解 Momentum 的关键:vt 是过去所有梯度的指数加权平均。
权重分布(β=0.9 时):
| 梯度 | 权重 |
|---|---|
| gt(最新) | 1 |
| gt−1 | 0.9 |
| gt−2 | 0.81 |
| ... | ↓ |
| gt−t′ | 0.9t′ |
有效窗口长度约 1−β1=10(β=0.9 时)。
2.2 为什么能抑制震荡
想象梯度序列 +1,−1,+1,−1,…:
- 无 Momentum:净位移 =(+1−1)×η=0,原地打转
- 有 Momentum(β=0.9):震荡被加权平均掉,v 收敛到一个非零的小正值,持续朝一个方向前进
本质:Momentum 是一个低通滤波器,把高频震荡滤掉,保留低频的漂移方向。
这就是为什么在陡峭峡谷(梯度方向来回变化)里,Momentum 能显著加速。
2.3 Nesterov Accelerated Momentum
改进:不用「当前」梯度算,而是用「位置更新后的」梯度。
θt=θt−1−η(βvt−1+gt(θt−ηvt−1))
直觉:先按动量往前走一步,再看那个位置的梯度(比原点更有信息量)。这叫「前瞻」。
在图像识别任务上比标准 Momentum 快约 5~10%。PyTorch 里对应 momentum=0.9, nesterov=True。
2.4 SGD 的实际地位
在视觉任务上,SGD+Momentum 至今仍然是首选(ResNet、BERT 训练都用 SGD),因为:
- 泛化性能通常略优于自适应方法
- 超参更少(只有 lr 和 momentum)
- 内存开销小
自适应方法(Adam 系)在 NLP / Transformer 上更常用,见第 4 节。
三、AdaGrad 与 RMSProp
3.1 AdaGrad:累加梯度平方
rt=rt−1+gt2,θt=θt−1−rtηgt
问题:rt 只增不减,分母越来越大 → lr 趋于 0,过早停止学习。
3.2 RMSProp:改成滑动平均
st=βst−1+(1−β)gt2,θt=θt−1−st+ϵηgt
两个改动:
- gt2 改成滑动平均 → st 可以增减,不再单调增
- 分母用 st(标准差的估计)而非 rt(平方和)
直觉:参数变动剧烈(梯度大)→ 除以大数 → 步长自动变小。每个参数获得自己的学习率。
这就是「自适应学习率」的核心思想。
四、Adam:动量 + 自适应学习率
4.1 公式
mtvtm^tv^tθt=β1mt−1+(1−β1)gt(一阶矩:动量)=β2vt−1+(1−β2)gt2(二阶矩:方差)=1−β1tmt(偏差修正)=1−β2tvt=θt−1−ηv^t+ϵm^t
4.2 两个状态的直觉
| 状态 | 统计什么 | 类比 |
|---|---|---|
| mt | 梯度的均值(方向) | 「最近梯度大致指向哪」 |
| vt | 梯度的平方均值(大小) | 「梯度波动多大」 |
更新规则 = 沿梯度的平均方向走,步长按波动幅度缩放。
更新量=η⋅尺度方向
如果某个参数的梯度一直是 5.0(方向一致、尺度大),更新量 ≈η⋅1,步长正常。 如果梯度是 +5,−5 交替(方向不一致),m≈0 → 更新量趋于 0,自动减小步长。
这就是 Adam 比 SGD 更快的原因:它自动识别并抑制震荡方向的更新。
4.3 偏差修正为什么必需(重要)
m0=v0=0,所以第一步:
m1=(1−β1)g1
这明显低估了梯度的真实均值(应该是 g1,但只得到了 0.1g1)。早期所有估计都偏向 0。
修正:除以 1−βt(因为 ∑i=0t−1(1−β)βi=1−βt,归一化因子):
| 步数 t | 1−β1t(β1=0.9) | 修正倍数 1/(1−βt) |
|---|---|---|
| 1 | 0.1 | 10× |
| 2 | 0.19 | 5.3× |
| 10 | 0.65 | 1.5× |
| 100 | 0.99997 | ≈1.0 |
所以 Adam 的前几步更新量被放大了最多 10 倍。 这个修正保证了 lr 在任何时刻都是「名义 lr」,不会被初始化效应干扰。
这也解释了 warmup 的一个理由:Adam 早期更新量本来就偏大(即使修正后),warmup 再进一步慢慢提 lr,能让早期训练更稳。
4.4 Adam 的致命缺陷:权重衰减不兼容
关键问题:weight decay 在 SGD 里是加到梯度上的:
θ←θ−η(∇L+λθ)=(θ−ηλθ)−η∇L
这个顺序很重要——先衰减,再走梯度。
但 Adam 里两者混在一起了:
θ←θ−η⋅v^t+ϵm^t
如果 weight decay 加进梯度 gt,它也会进入 mt 和 vt:
mt=β1mt−1+(1−β1)(∇L+λθ)
后果:
- 衰减量被 v^t1 缩放。如果某个参数的梯度一直很小(v^t 很小),衰减量会被放大——本来想轻微收缩,结果剧烈收缩
- 衰减量与梯度历史耦合,无法控制
这就是 AdamW(Adam with decoupled weight decay)的动机:把衰减从梯度里拿出来,直接作用于参数:
θt=θt−1−ηv^t+ϵm^t−ηλθt−1
衰减项不再进入自适应缩放,永远是恒定的 ηλθ。
4.5 AdamW 的额外好处:可以配合 decoupled 之外的策略
因为解耦了,衰减项可以是任何形式,而不只是 L2:
- decoupled weight decay(AdamW 本身)
- Layer-wise / 部分衰减:某些层不衰减(比如最后一层)
- 与学习率调度正交:cosine schedule 下两者互不干扰
五、优化器选择决策树
任务类型?
├─ 图像分类 / 视觉 CNN
│ └─ SGD + Momentum(0.9) + Nesterov ← 泛化好,验证集指标通常更高
│
├─ Transformer / NLP 微调
│ └─ AdamW(lr=1e-4~3e-4, betas=(0.9,0.95), wd=0.01~0.1)
│ ↑ 注意 β2 常设0.95 而非 0.999
│
├─ LLM 预训练
│ └─ AdamW + cosine schedule + warmup
│ lr=1e-4~3e-4,wd=0.1,grad clip=1.0
│
└─ 小模型 / 实验 / 数据量极少
└─ Adam 系更宽容,不用精调 lr2
3
4
5
6
7
8
9
10
11
12
13
14
为什么 Transformer 用 AdamW 而不是 SGD?
- 各种尺度的注意力参数需要不同 lr,SGD 处理不好
- 稀疏梯度(部分维度梯度为0)下Adam 更稳
- Transformer 训练对小 lr 很敏感,Adam 的自适应让调参更容易
为什么视觉任务仍偏爱 SGD?
一个被广泛引用的经验:AdamW 类方法在训练后期会明显过拟合验证集,而 SGD+momentum 泛化更好。可能的解释是自适应方法的有效步长不稳定。文献很多、结论不完全统一,实践中以验证集结果为准。
六、β2 的一个实用细节
默认值 β2=0.999 在 LLM 训练中通常改成 0.95。 为什么?
1−β21 = 有效窗口长度:
| β2 | 窗口长度 |
|---|---|
| 0.999 | 1000 步 |
| 0.95 | 20 步 |
| 0.99 | 100 步 |
LLM 训练中 β2=0.999 的窗口太长:因为训练过程中梯度分布本身在快速变化(从初期的大幅下降到后期的精细),用 1000 步的历史平均会让 vt 严重滞后于当前状态。
改成 0.95(20 步窗口)后,二阶矩能更快跟上梯度的当前尺度。这是 HuggingFace、LLaMA 训练脚本的默认配置,也是现在的事实标准。
七、动手实验
实验 1:亲手实现 SGD / Momentum / Adam
比调torch.optim 有价值得多——你会真正理解每个优化器的「状态」是什么。
import torch
def sgd(p, g, lr=0.1):
p -= lr * g
def momentum(p, g, lr=0.1, beta=0.9, v=None):
if v is None: v = torch.zeros_like(p) # 状态:速度
v.mul_(beta).add_(g)
p -= lr * v
return v
def adam(p, g, lr=0.1, b1=0.9, b2=0.999, eps=1e-8, m=None, v=None, t=0):
if m is None: m, v, t = torch.zeros_like(p), torch.zeros_like(p), 0
t += 1
m = b1*m + (1-b1)*g # 状态:一阶矩
v = b2*v + (1-b2)*g**2 # 状态:二阶矩
mh, vh = m/(1-b1**t), v/(1-b2**t) # 偏差修正
p -= lr * mh / (vh.sqrt() + eps)
return m, v, t
# 交替震荡的梯度序列,模拟"峡谷"地形(梯度方向来回变)
grads = [torch.tensor([1.0 if i % 2 == 0 else -1.0]) for i in range(10)]
for name in ["SGD", "Momentum", "Adam"]:
p = torch.tensor([1.0])
state = {}
for g in grads:
if name == "SGD":
sgd(p, g, lr=0.1)
elif name == "Momentum":
state['v'] = momentum(p, g, lr=0.1, v=state.get('v'))
else:
state['m'], state['v'], state['t'] = adam(
p, g, lr=0.1, m=state.get('m'), v=state.get('v'), t=state.get('t', 0))
print(f"{name:>10}: 起点 1.0 → 最终 {p.item():+.4f}")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
实测输出:
SGD: 起点 1.0 → 最终 +1.0000 ← ★ 净位移为 0,完全原地打转
Momentum: 起点 1.0 → 最终 +0.6915 ← 动量把它推过去了
Adam: 起点 1.0 → 最终 +0.8455 ← 自适应缩放,步子更小但也有效2
3
SGD 最终精确回到 1.0,因为 10 个梯度里 5 个 +1、5 个 -1,净和为 0。这就是震荡梯度下朴素 SGD 的困境——它把力气全花在互相抵消上了。
Momentum 和 Adam 都成功偏离了原点,因为它们的状态变量记住了历史,震荡被平均掉。
踩坑记录:我第一版写成
zip(params, grads)的批量版,结果三个优化器输出完全一样(都是 0.9)。原因是我只传了1 个参数却给了 10 个梯度,zip按最短长度截断,等于只做了1 步。单个参数 + 一串梯度这种场景,必须写成逐次调用并显式传递状态,不能用批量接口。
实验 2:验证 Adam 的偏差修正
import torch
print("=== 第一步:不做偏差修正 vs 做修正 ===")
for use_bias_correction in [False, True]:
torch.manual_seed(0)
p = torch.tensor([0.0])
m = torch.zeros(1); v = torch.zeros(1)
g = torch.tensor([3.0]) # 固定的梯度
lr, b1, b2 = 0.1, 0.9, 0.999
m = b1*m + (1-b1)*g
v = b2*v + (1-b2)*g**2
if use_bias_correction:
mh, vh = m/(1-b1**1), v/(1-b2**1)
else:
mh, vh = m, v
step = lr * mh / (vh.sqrt() + 1e-8)
print(f"bias_correction={use_bias_correction}: m={m.item():.4f} "
f"更新量={step.item():.4f}")
print("\n理论上:|g|=3 时,Adam 的步长应该是 lr=0.1 才对")
print("不做修正 -> 步长被缩小到 lr*0.1 = 0.01(因为 m 只积累到 0.3)")
print("做了修正 -> 步长 ≈ 0.1(正确)")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
实测输出:
corr=False: step=0.3162 ← 比理论值 0.1 大3 倍
corr=True: step=0.1000 ← ★ 精确等于 lr2
注意:不做修正时步长是 0.3162,比正确值 0.1 大了 3 倍,不是「缩小」。
为什么反直觉? 因为 ϵ 也在起作用。看清楚 β2=0.999 时发生了什么:
| 修正项 | 值 |
|---|---|
| m1 | 0.1×3=0.3 |
| v1(未修正) | 0.001×9=0.009 |
| v1(修正后) | 0.009/(1−0.999)=9.0 |
v1未修正=0.0949,v1修正=3.0
分子 m^ 从 0.3 → 3.0(放大 10 倍),分母也从 0.095 → 3.0(放大约 31 倍)。分母放得更多,所以最终步长偏小——0.1 × 0.3/0.095 = 0.316,正是实测值。
一句话总结偏差修正的作用:让第一步的有效步长精确等于名义 lr,不被初始化时的 β 衰减拖慢。
实验 3:Adam vs AdamW 的差异(理解解耦的价值)
import torch
print("=== 一个梯度极小的参数,weight decay 会怎样? ===")
torch.manual_seed(0)
for name, use_coupled in [("Adam (耦合,衰减进梯度)", True), ("AdamW (解耦)", False)]:
p = torch.tensor([1.0], requires_grad=True)
m, v = torch.tensor([0.0]), torch.tensor([0.0])
lr, b1, b2, eps, wd = 0.1, 0.9, 0.999, 1e-8, 0.1
# 连续 3 步,梯度恒为 0.01(很小)
for t in range(1, 4):
g = torch.tensor([0.01])
if use_coupled:
g_eff = g + wd * p.detach() # 衰减加进梯度
else:
g_eff = g
m = b1*m + (1-b1)*g_eff
v = b2*v + (1-b2)*g_eff**2
mh, vh = m/(1-b1**t), v/(1-b2**t)
p.data -= lr * mh / (vh.sqrt()+eps)
if not use_coupled:
p.data -= lr * wd * p.data # 解耦:直接衰减
print(f"{name:>28}: p={p.item():.6f}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
实测输出:
Adam (耦合,衰减进梯度): p=0.701382
AdamW (解耦):p=0.6762592
梯度恒为 0.01(很小),weight decay=0.1,目标是把参数拉向 0:
- AdamW:衰减量精确是 ηλp=0.1×0.1×p,行为完全可预测
- Adam:衰减量被送进 vt,最终的步长 ∝v^t1 包含了衰减的历史,于是衰减强度和梯度历史纠缠在一起
数值上这里只差 3%,看起来不多。但在真实训练里差异会累积:Adam 的 vt 被污染后,会影响所有参数的自适应步长,而不只是被衰减的那些——这才是问题的严重性所在。
想看更戏剧性的差异,把梯度再调小试试(v^t 越小,1/v^t 的放大效应越强)。
实验 4:β2 的影响
import torch
print("=== β2 决定的二阶矩窗口长度 ===")
for b2 in [0.9, 0.95, 0.99, 0.999]:
print(f"β2={b2:<6} 有效窗口 = 1/(1-β2) = {1/(1-b2):.0f} 步")
print("\nLLM 预训练默认 β2=0.95(20 步窗口)而非 0.999(1000 步)")
print("原因:训练中梯度尺度快速变化,需要二阶矩跟上当前状态")2
3
4
5
6
八、自测题
Q1:Adam 的 ϵ 是干什么的?如果设成 0 会有什么问题?
答案
防止除零。当某个参数在整个训练过程中梯度都是 0(可能因为它被 mask 了、或者初始化为 0 后从未更新),vt=0,则 vtmt 是 0/0。
加了 ϵ(默认 1e-8)后变成 vt+ϵmt,分母有下界。
设成 0 的风险:
- 分母为 0 → NaN
- 即使不为 0,ϵ 太小会放大数值误差
- 实践中 ϵ 的选择对结果影响很小,因为它加在 vt 上(vt 通常远大于 1e-16)
Q2:为什么 SGD 在图像任务上泛化性常优于 Adam?
答案
没有公认定论,文献很多,主要假说:
- 自适应方法的等效学习率不稳定。Adam 的 η/vt 在梯度尺度变化时会剧烈波动,导致它找到的是「容易优化的解」而非「泛化好的解」(类似 sharpness-aware minimization 的视角)。
- Adam 的隐式正则。m/v 相当于对梯度做归一化,抹平了不同参数方向的尺度差异,导致模型倾向于找到「平坦但可能不泛化」的方向。
- 训练后期行为差异。Adam 在训练后期有效步长不稳定,容易在验证集上过拟合。
但这个结论不绝对:
- 有研究(e.g. 在部分视觉 benchmark 上)发现 AdamW 配合好的调度也能超过 SGD
- ConvNeXt 等现代视觉模型用 AdamW 也拿到了 SOTA
实践建议:视觉任务两个都试,用验证集决定。不要预设。
Q3:一个参数的梯度一直是 0.001,另一个是 100。用 SGD 和 Adam 各走一步,哪个参数动得多?
答案
SGD(lr=0.1):两个参数的更新量分别是 0.0001 和 10。梯度大的动得多,可能直接发散。
Adam:
- 参数A(梯度一直 0.001):m→0.001,v→0.001,更新量 =0.1×0.0010.001=0.1
- 参数B(梯度一直 100):m→100,v→100,更新量 =0.1×100100=0.1
两个参数走一样多! 这就是自适应学习率的意义——不管梯度绝对大小如何,每个参数每步都走大致相同的距离。
这个特性对 Transformer 特别重要:注意力层的 Q/K/V 矩阵、FFN 的两层,梯度尺度差异巨大(可能相差几个数量级)。SGD 处理不好,Adam 能自动处理。
Q4:为什么 LLM 都用 warmup?如果没有 warmup 会怎样?
答案
warmup 的作用:在训练最初的几百/几千步把 lr 从 0 线性提升到目标值。
没有 warmup 的问题:
初期梯度不稳。训练开始时参数还没进入「有意义的空间」,梯度方向噪声大。直接用大 lr 会把参数推到糟糕的区域,一旦模型已经学到的结构被破坏,后面很难恢复。
自适应优化的二阶矩估计不准。mt,vt 在头几步是严重有偏的(虽然有偏差修正,但修正本身让前几步的有效 lr 偏大),此时用满 lr 不稳。
配合 AdamW 会更糟。因为 weight decay 的实际强度是 ηλθ——如果 lr 一开始就是满值,衰减也是满的,参数会在最初几步被剧烈拉向 0。
warmup 之后:
- 梯度分布稳定了
- 二阶矩估计准确了
- 再配合 cosine decay,让 lr 在训练后期精细收敛
LLaMA 的做法:warmup 2000 步 + cosine decay到 10% 峰值。这套配置几乎成了标准。
Q5:PyTorch 里 Adam(lr=..., weight_decay=...) 和 AdamW 到底差在哪?如果我用了 Adam + weight_decay,我的模型会有什么问题?
答案
差在weight decay 的施加位置:
# Adam(L2 正则耦合)
g_eff = g + weight_decay * p # 衰减进梯度
m = b1*m + (1-b1)*g_eff # 衰减进动量
v = b2*v + (1-b2)*g_eff**2 # 衰减进二阶矩 ← 问题在这
p -= lr * m_hat / (sqrt(v_hat)+eps)
# AdamW(解耦)
g_eff = g # 梯度不含衰减
m, v = ...(只用数据梯度)...
p -= lr * m_hat / (sqrt(v_hat)+eps)
p -= lr * weight_decay * p # 衰减独立施加2
3
4
5
6
7
8
9
10
11
实际问题:
- 衰减项被 v^1 缩放。梯度历史小的参数,衰减被放大——本意轻微收缩,实际剧烈收缩
- 衰减量和梯度历史纠缠,无法单独控制强度
- 在 warmup 阶段(lr 变化)两者行为差异更大
实际影响:在很多任务上 Adam + weight_decay 也能训出好模型,所以这个bug 不会立刻暴露。但在大规模预训练 + 长训练里,AdamW 的稳定优势才显现出来。这就是为什么 LLM 全用 AdamW。
建议:直接用 AdamW。没有理由用 Adam + weight_decay。
下一篇 → 反向传播的完整推导
Part 2 · 深度学习原理
这一部分讲深度网络特有的机制。为什么深网络能工作、为什么会出现梯度消失、以及那些「魔法数字」背后是什么。
| # | 章节 | 核心问题 |
|---|---|---|
| 05 | 反向传播的完整推导 | 链式法则在这个网络里怎么具体展开?梯度消失的本质是什么? |
| 06 | 归一化:从 BatchNorm 到 RMSNorm | 归一化到底在解决什么?为什么 BatchNorm 淘汰而 LayerNorm 统治? |
| 07 | 卷积与感受野 | 卷积为什么有效?什么条件下等价于全连接? |
| 08 | 残差连接与网络架构 | 退化问题的数学本质是什么? |
| 09 | 初始化与数值稳定 | 为什么不能用 0 初始化?为什么深层需要特定初始化? |
这一部分要建立的判断力
- 训练不收敛时,能按顺序排查:初始化 → 归一化 → 残差 → 学习率
- 看到一个「训练技巧」时,能说清它在数学上改变了什么
- 能区分优化困难和表达能力不足 —— 这两个的解法完全不同
第 5 篇是硬骨头
它把反向传播在一个具体网络上一项一项推完。第一次读会慢,但推完一遍之后, 第 6-9 篇里所有「为什么这样设计」都会变得显然。
5 · 反向传播的完整推导
核心问题:链式法则在一个具体网络里怎么展开?梯度消失/爆炸的数学本质是什么?
一、一个最小的网络
用最小的例子把整个过程走一遍。网络只有两层:
x --W1--> z1 --ReLU--> a1 --W2--> z2 --> L
+b1 +b22
前向:
z1a1z2L=W1x+b1=ReLU(z1)=max(0,z1)=W2a1+b2=(z2−y)2
现在要算 ∂W1∂L、∂W2∂L、∂b1∂L、∂b2∂L。
这里的关键认知:反向传播不是「每个参数独立算一次」,而是从输出往输入传一个「敏感度」,沿途用乘法法则分发给每个分支。
二、反向传播:一个参数一个参数推
2.1 从损失开始
∂z2∂L=2(z2−y)
这个 2 就是 MSE 的特征。如果 loss 写成 21(z2−y)2(很多教材这么写),梯度就干净了:∂z2∂L=(z2−y)。这就是为什么有些教材的 loss 系数是 2N1 而不是 N1 ——纯粹为了消掉这个 2。
2.2 分发给 W2 和 b2
z2=W2a1+b2,求 ∂W2∂L:
z2=j∑(W2)ij(a1)j+(b2)i
对 (W2)ij 求导,(W2)ij 只出现在第 i 项里:
∂(W2)ij∂L=∂z2∂L⋅∂(W2)ij∂z2=∂z2∂L⋅(a1)j
写成矩阵形式(避免下标地狱):
∂W2∂L=[1,1]∂z2∂L⋅[1,H]a1⊤=outer product
外积的形状:[out,1]×[1,in]=[out,in] ✓ 正好是 W2 的形状。
∂b2∂L=∂z2∂L
(因为 b2 是加法,导数为 1)
2.3 穿过 ReLU
∂z1∂a1={10z1>0z1≤0
所以:
∂z1∂L=∂a1∂L⊙1[z1>0]
这里有个重要事实:ReLU 会把负半轴的梯度完全掐断(置0)。 这是 ReLU 「死神经元」问题的根源——如果一个神经元对所有输入都是负的,它的梯度永远是 0,参数永远不更新,等于废掉了。
2.4 分发给 W1 和 b1
∂W1∂L=∂z1∂L⋅x⊤,∂b1∂L=∂z1∂L
到这里就完成了。 完整链条:
∂z2∂L→∂a1∂L→∂z1∂L→∂W1∂L,∂b1∂L
三、通用规则(记住这三条就够)
反向传播就是从后往前反复应用这三条规则:
规则 1:加法 → 梯度直接分发
z = a + b
∂L/∂a = ∂L/∂z ∂L/∂b = ∂L/∂z2
多个分支收到相同的上游梯度(所以 GradienT 累加不是 bug,是数学要求)。
规则 2:乘法 → 梯度交叉相乘
z = a × b
∂L/∂a = ∂L/∂z × b ∂L/∂b = ∂L/∂z × a2
规则 3:矩阵乘法 → 外积
z = a @ W (a: [B,in], W: [in,out], z: [B,out])
∂L/∂W = aᵀ @ (∂L/∂z) [in,B]@[B,out] = [in,out] ✓ 和 W 同形状
∂L/∂a = (∂L/∂z) @ Wᵀ [B,out]@[out,in] = [B,in] ✓ 和 a 同形状
∂L/∂b = (∂L/∂z).sum(0) 按 batch 求和→ [out]2
3
4
5
记忆方法:「哪个输入和梯度同形状,就用它的转置去乘另一个」。
- W 是
[in, out],要得到[in, out],就用 a⊤[B,in] 乘 (∂L/∂z)[B,out] - a 是
[B, in],要得到[B, in],就用 (∂L/∂z)[B,out] 乘 W⊤[out,in]
这就是外积的批量版:∂L/∂W=∑i(第i 个样本的 ∂L/∂zi)⊗ai —— 每个样本贡献一个外积,加起来就是矩阵乘。
PyTorch 里的对应关系:
| 数学写法 | PyTorch |
|---|---|
| 外积 (ab⊤) | torch.outer(a, b) 或 a @ b.T |
| 张量对元素积 ⊙ | a * b |
| 矩阵乘 | @ |
| 沿 dim 求和 | .sum(dim=n) |
一句话总结
反向传播 = 反向应用链式法则 + 把「敏感度」沿计算图分发。 加法节点等值分发,乘法节点交叉相乘,矩阵乘用外积。
四、梯度消失与爆炸的数学本质
这是本文最重要的部分。理解了这段,你就能自己判断任何架构设计的梯度行为,不需要背「ResNet 解决了梯度消失」这种结论。
4.1 梯度经过一层的衰减/放大
考虑一个 L 层、宽度 n 的全连接网络(ReLU + 合适的初始化)。每层权重的 Jacobian 矩阵 Ji 的元素典型大小约 O(n1)(He 初始化的设计目标)。
反向传播时梯度要连乘 L 个 Jacobian:
∂x∂L=JL⋅JL−1⋯J1⋅∂out∂L
每个因子贡献一个 n1,连乘 L 次:
(n1)L=n−L/2
| 网络 | L | n | n−L/2 |
|---|---|---|---|
| 浅层 | 3 | 100 | 10−3 |
| 深层 | 30 | 100 | 10−15 |
| 深层 | 100 | 1000 | 10−150 |
10−15 在 float32 的精度下就是 0。 这就是梯度消失。
4.2 Sigmoid 的额外问题
如果激活函数是 sigmoid/tanh,导数最大只有 0.25(sigmoid):
σ′(z)=σ(z)(1−σ(z))≤0.25
每经过一层,梯度至少乘以 0.25。 10 层就是 0.2510≈10−6。
为什么 ReLU 缓解了这个问题:ReLU′(z)≥0 且期望为 1(对 z>0 部分导数为 1),所以不会引入额外的衰减。这就是 2012 年 ReLU 带来突破的数学原因。
4.3 梯度爆炸
反过来,如果权重初始化太大(J 的元素 ≫1),连乘后会指数爆炸:
∂x∂L∝(nsomething≫1)L→∞
表现:loss 突然变成 nan、梯度值离谱、参数被更新到溢出。
解法是梯度裁剪(gradient clipping):
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)把所有梯度作为一个整体,计算 L2 范数,如果超过阈值就等比缩放:
g^=g⋅∥g∥max_norm当∥g∥>max_norm
注意是等比缩放所有梯度,不是逐元素裁剪(后者会改变梯度方向,破坏优化语义)。
LLM 训练标准配置里都有 clip=1.0。这不是可选项。
4.4 残差连接如何解决
残差连接:al+1=al+f(al)
反向传播时,梯度有一条「直达通路」:
∂al∂L=∂al+1∂L(I+∂al∂f)
那个 I(单位矩阵)意味着梯度可以原封不动地传下去。
即使 f 那部分的梯度衰减到 0,还有 I 兜着。所以:
∂a1∂L=l=1∏L(I+Jl)
这个乘积里,每项都含 I,展开后至少有一项全是 I 的乘积(对应「所有残差路径都不走 f」那条路线),所以梯度不衰减。
这是 ResNet 能训 100 层的数学保证。 详见第 8 篇。
4.5 一个必须掌握的自测方法
判断一个新架构会不会梯度消失/爆炸,不要靠猜——做实验测量。
# 逐层测量梯度范数,一眼看出是否衰减/爆炸
for name, param in model.named_parameters():
if param.grad is not None:
print(f"{name:40s} ||grad|| = {param.grad.norm().item():.3e}")2
3
4
判读:
- 逐层指数下降(如 1e-1 → 1e-3 → 1e-5 → 1e-8)→ 梯度消失
- 逐层指数上升 → 梯度爆炸
- 各层量级相当(1e-2 量级上下浮动)→ 健康
这个技巧在任何论文的复现里都用得上。看到一个新架构,第一件事就是打这个表。
五、动手实验
实验 1:验证手推公式和 autograd 完全一致
import torch
torch.manual_seed(0)
B = 4
W1 = torch.randn(3, 2, requires_grad=True); b1 = torch.zeros(2, requires_grad=True)
W2 = torch.randn(2, 1, requires_grad=True); b2 = torch.zeros(1, requires_grad=True)
x = torch.randn(B, 3); y = torch.randn(B, 1)
# ---- 前向 ----
z1 = x @ W1 + b1 # [4,3]@[3,2] = [4,2]
a1 = torch.relu(z1)
z2 = a1 @ W2 + b2 # [4,2]@[2,1] = [4,1]
loss = ((z2 - y) ** 2).mean()
loss.backward()
# ---- 手推反向(严格按矩阵形状推导)----
# L = (1/B) * Σ_i (z2_i - y_i)²
# dL/dz2 是逐样本的:∂L/∂z2_i = 2(z2_i - y_i)/B
dL_dz2 = 2 * (z2 - y) / B # [4,1]
# W2 是 [2,1],z2 = a1 @ W2 + b2 → dL/dW2 = (dL/dz2)ᵀ @ a1,但需要 [2,1]
dL_dW2 = dL_dz2.T @ a1 # [1,4]@[4,2] = [1,2]
dL_db2 = dL_dz2.sum(0) # [1]
# z2 = a1 @ W2 + b2 → dL/da1 = dL/dz2 @ W2ᵀ
dL_da1 = dL_dz2 @ W2.T # [4,1]@[1,2] = [4,2]
# ReLU:梯度在负半轴为 0
dL_dz1 = dL_da1 * (z1 > 0).float()# [4,2]
# z1 = x @ W1 + b1,W1 是 [3,2]
dL_dW1 = x.T @ dL_dz1 # [3,4]@[4,2] = [3,2]
dL_db1 = dL_dz1.sum(0) # [2]
print("形状对齐(必须和参数完全一致):")
print(f" dL_dW1 {tuple(dL_dW1.shape)} vs W1 {tuple(W1.shape)}")
print(f" dL_db1 {tuple(dL_db1.shape)} vs b1 {tuple(b1.shape)}")
print(f" dL_dW2.T {tuple(dL_dW2.T.shape)} vs W2 {tuple(W2.shape)}")
print(f" dL_db2 {tuple(dL_db2.shape)} vs b2 {tuple(b2.shape)}")
print("\n与 autograd 对比(注意 W2 需要转置):")
pairs = [("dL_dW1", dL_dW1, W1), ("dL_db1", dL_db1, b1),
("dL_dW2", dL_dW2.T, W2), ("dL_db2", dL_db2, b2)]
for name, mine, param in pairs:
err = (mine - param.grad).abs().max().item()
print(f" {name}: 最大误差 = {err:.1e}")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
实测输出:
形状对齐(必须和参数完全一致):
dL_dW1 (3, 2) vs W1 (3, 2)
dL_db1 (2,) vs b1 (2,)
dL_dW2.T (2, 1) vs W2 (2, 1)
dL_db2 (1,) vs b2 (1,)
与 autograd 对比(注意 W2 需要转置):
dL_dW1: 最大误差 = 0.0e+00
dL_db1: 最大误差 = 0.0e+00
dL_dW2: 最大误差 = 0.0e+00
dL_db2: 最大误差 = 0.0e+002
3
4
5
6
7
8
9
10
11
误差全是 0(浮点运算顺序一致时)。这证明了你的推导和 autograd 做的是同一件事。
⚠️ 手推时我踩的三个坑(重要)
这个实验的价值一半在于踩坑。三个错误都是我第一次写时犯的:
坑 1:忘了 batch 平均的系数
loss = (z2-y)**2.mean() 是对 batch 和输出维都求平均,所以:
∂z2∂L=B2(z2−y)
我一开始写成 2*(z2-y),漏了除以 B=4,结果所有梯度差 4 倍。
坑 2:矩阵乘法顺序写反
z1=x@W1 中 W1 是 [3,2](in×out),所以:
∂W1∂L=x⊤⋅∂z1∂L(先转 x)
我写成了 dL_dz1.T @ x,结果形状是 [2,3] 而不是 [3,2]。
PyTorch 会静默地广播出一个错误结果而不是报错,所以这个 bug 特别危险——必须打印形状。
坑 3:torch.outer 的参数必须是 1维
torch.outer(dL_dz2, a1) 报错,因为 dL_dz2 是 [4,1] 而不是 [1]。
修正写法是 dL_dz2.T @ a1,等价于外积但对形状没要求。
教训:推导和代码之间隔着一层形状地狱。写完推导第一件事是打印 .shape 和参数的 .shape 对比,这比什么都重要。
实验 2:测量梯度消失
import torch
from torch import nn
def measure_grad_norms(depth, activation='relu'):
layers, ins = [], 20
act = nn.ReLU if activation == 'relu' else nn.Sigmoid
for _ in range(depth):
layers += [nn.Linear(ins, 20), act()]
ins = 20
layers += [nn.Linear(20, 1)]
model = nn.Sequential(*layers)
x = torch.randn(64, 20)
for p in model.parameters():
p.grad = None
model(x).sum().backward()
norms = [p.grad.norm().item() for p in model.parameters()]
return norms
print("=== ReLU 网络:逐层梯度范数 ===")
for depth in [2, 6, 20]:
norms = measure_grad_norms(depth)
print(f"\n深度 {depth}: 首层 {norms[0]:.2e} → 末层 {norms[-1]:.2e}"
f" 衰减倍数 {norms[0]/norms[-1]:.1f}x")
print("\n=== Sigmoid 网络 ===")
for depth in [2, 6, 20]:
norms = measure_grad_norms(depth, 'sigmoid')
print(f"深度 {depth}: 首层 {norms[0]:.2e} → 末层 {norms[-1]:.2e}"
f" 衰减倍数 {norms[0]/norms[-1]:.1f}x")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
你会观察到:ReLU 网络的梯度衰减远慢于 Sigmoid;层数增加时两边都衰减,但 Sigmoid 更严重。这就是 ReLU 优于 Sigmoid 的定量证据。
实验 3:残差连接如何消除衰减
import torch
from torch import nn
class Plain(nn.Module):
def __init__(self, depth, width=20):
super().__init__()
self.layers = nn.ModuleList([nn.Linear(width, width) for _ in range(depth)])
def forward(self, x):
for l in self.layers: x = torch.relu(l(x))
return x.sum()
class Res(nn.Module):
def __init__(self, depth, width=20):
super().__init__()
self.layers = nn.ModuleList([nn.Linear(width, width) for _ in range(depth)])
def forward(self, x):
for l in self.layers:
x = torch.relu(l(x) + x) # ★ 残差连接
return x.sum()
print(f"{'深度':>6} {'普通网络衰减':>14} {'残差网络衰减':>14}")
for depth in [2, 6, 20, 50]:
x = torch.randn(64, 20)
out = []
for M in [Plain, Res]:
torch.manual_seed(42); m = M(depth)
for p in m.parameters(): p.grad = None
m(x).backward()
n = [p.grad.norm().item() for p in m.parameters()]
out.append(n[0] / n[-1])
print(f"{depth:>6} {out[0]:>13.1f}x {out[1]:>13.1f}x")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
你会看到:深度越大,普通网络的梯度衰减倍数呈指数增长,而残差网络几乎保持恒定。这就是 ResNet 能训 100+ 层的定量证据。
六、自测题
Q1:为什么 ReLU 会出现「死神经元」?什么条件下会出现?
答案
机制:ReLU 的梯度是 1[z>0]。如果某个神经元对所有输入都输出 z≤0,它的梯度永远是 0,参数永远不更新。
什么条件会触发:
- 初始化不当:权重初始化为全 0 或负偏置过大,初始输出就大面积为负
- 学习率过大:一次更新把参数推到负半轴,之后再也回不来
- 输入分布本身为负(比如前面接了负偏置、或者数据本身偏负)
为什么回不来:一旦某个神经元对所有输入都是负输出,它就不会参与任何有用的计算,其他神经元也不会因为它而得到有用的梯度(梯度是累加的,但这个神经元贡献 0),所以整个系统失去了这个神经元。这就是「死」。
解法:
- Leaky ReLU:max(0.05z,z),给负半轴一个小的非零斜率
- 参数化 ReLU(PReLU):斜率可学习
- 正确初始化(He 初始化,让初始输出正负各半)
- 合适的 lr
Q2:如果网络所有权重都初始化为 0,会发生什么?
答案
输出恒为 0,且梯度也为 0,训练完全不进行。
推导:
- 前向:z=Wx+b,若 W=0 则 z=0(不管 b 多大),所有层输出 0
- 反向:∂W2∂L=∂z2∂L⋅a1⊤,而 a1=0,所以 ∂W2∂L=0
- 逐层回推,所有梯度都是 0
对称性问题:如果所有神经元用同一个初始化且输入也一样,它们的梯度完全相同 → 参数更新也相同 → 永远保持对称,等效于一个神经元。
注意区分:
- 权重 W 不能初始化为 0
- 偏置 b 可以初始化为 0(因为它不造成对称性问题)
- 最后一个全连接层的 bias 可以初始化为 0,但它的 weight 初始化为 0 会导致 logits 全相同 → 初始 loss 就是 logK
这也是为什么 PyTorch 里的 nn.Linear 默认 bias=True 且初始化均匀。
Q3:梯度裁剪是把超过阈值的部分「截断」,为什么等比缩放(全梯度乘一个系数)更好?
答案
逐元素裁剪(clip by value):
gi←max(−ϵ,min(ϵ,gi))
这会改变梯度的方向。比如原梯度是 (0.1,10.0),裁剪后是 (0.1,1.0),方向从近似 y 轴变成更接近 x 轴。优化的方向被改掉了——这违反了「梯度下降沿最陡下降方向」的前提。
等比缩放(clip by norm):
g^=g⋅∥g∥ϵ
所有分量乘同一个系数,方向完全不变,只是整体步长被限制。
实现:
torch.nn.utils.clip_grad_norm_(
model.parameters(),
max_norm=1.0,
error_if_nonfinite=False, # nan 时跳过
)2
3
4
5
顺便解释一下 clip 的另一个作用:当出现 nan(数值爆炸的极端情况)时,clip_grad_norm_ 遇到 nan 会把整个梯度设为 0(见源码),相当于「跳过这一步」,避免 nan 污染参数。这给了训练一层保护。
Q4:一个网络在前 10 层梯度正常,第 50 层开始梯度极小。用什么手段能诊断出根本原因?
答案
先做逐层梯度测量(实验 2 的方法),定位衰减从哪一层开始。
然后分类可能原因:
| 原因 | 判据 | 解法 |
|---|---|---|
| 深层 tanh/sigmoid | 衰减从第一个非线性层就开始 | 换 ReLU/GELU |
| 初始化太小 | 各层梯度都在衰减,且权重值本身很小 | 换 He 初始化 |
| 权重衰减过强 | 权重值被压得很小 → 梯度 = W · (上游) 也很小 | 调小 weight_decay |
| 残差路径缺失 | 纯串行结构,理论上就衰减 | 加残差连接 |
| 没有归一化 | 深层训练慢且梯度不稳 | 加 LayerNorm |
最快的验证方法:做一个受控实验。取一个纯 ReLU + He 初始化的深层 MLP,逐层打印梯度。如果这个基线是健康的,说明问题在你的具体设计(初始化、归一化、结构);如果基线也不健康,那就是深度本身带来的,需要残差连接。
# 最小复现
model = nn.Sequential(*[nn.Sequential(nn.Linear(64,64), nn.ReLU()) for _ in range(50)])2
6 · 归一化:BatchNorm → RMSNorm
核心问题:归一化到底在解决什么?为什么 BatchNorm 在 CNN 里统治、在 NLP 里被淘汰?RMSNorm 又省了什么?
一、归一化解决的两个问题
1.1 尺度问题(forward)
考虑一个深网络里的激活值。如果第 l 层的输出尺度是第 l−1 层的 k 倍(k>1),那么到第 L 层尺度就是 kL。
k=1.1、L=50 时,kL≈117。指数增长,一层层放大,网络极易饱和(sigmoid 饱和区 / ReLU 全死)。
归一化把每层的输出强制拉到「零均值单位方差」,尺度被控制住了。
1.2 优化问题(backward)
梯度里含因子 W⊤。如果 W 的谱半径(最大奇异值)ρ(W)>1,梯度连乘会指数爆炸;ρ(W)<1 则指数消失。
归一化让每层的 Jacobian 更接近「谱半径为 1 的等距映射」,梯度既不爆炸也不消失。
归一化的核心价值:让每一层的梯度尺度可控,从而让深层网络可训练。
二、BatchNorm 详解
2.1 公式
对一个 mini-batch 的同一通道,统计均值和方差,归一化后再缩放平移:
x^i=σB2+ϵxi−μB,yi=γx^i+β
其中(对 batch 维度 m 求):
μB=m1i∑xi,σB2=m1i∑(xi−μB)2
γ,β 是可学习的(shape = 通道数),作用是「让网络可以自己决定要不要归一化、以及归一化到什么分布」。如果 γ=β=0,BN 就退化成恒等映射——所以加 BN 的位置总会初始化成 γ=1,β=0。
2.2 关键设计:训练/推理行为不同
推理时没有 batch 可用,怎么办?
用训练时积累的滑动平均(running mean / running var):
# PyTorch 内部维护
self.register_buffer('running_mean', torch.zeros(num_features))
self.register_buffer('running_var', torch.ones(num_features))2
3
训练时更新:
running_mean=(1−momentum)⋅running_mean+momentum⋅μB
2.3 推理时必须 model.eval()(否则结果乱飞)
这是 BN 最容易踩的坑,也是第 2 篇强调过的 bug:
model.train():用当前 batch 的统计量model.eval():用训练时积累的 running 统计量
同一个输入,train 模式下每次预测结果都不同;eval 模式下才确定。
这就是为什么 LLM 普遍不用 BatchNorm:NLP 任务 batch 通常只有 1~8(序列长、显存紧),batch 统计量噪声极大,训练不稳定。
三、LayerNorm:Transformer 的选择
3.1 与 BN 的关键差异
| BatchNorm | LayerNorm | |
|---|---|---|
| 统计维度 | batch 维 + 空间维(对整个 batch 统计) | 特征维(每个样本独立算) |
| 依赖 batch 大小 | 是 | 否 |
| 训练/推理行为不同 | 是 | 否 |
| 适合序列 | 否(batch小 + 变长序列) | 是 |
| 消融论文 | ResNet | Transformer / LLaMA |
3.2 公式
对单个样本的特征维度统计:
μ=d1j=1∑dxj,σ2=d1j=1∑d(xj−μ)2
x^=σ2+ϵx−μ,y=γ⊙x^+β
注意:μ,σ 是对单个样本的所有特征算的,和其他样本无关。
3.3 为什么 NLP 必须用 LayerNorm
三个决定性原因:
- batch 统计量不可靠。NLP 的 batch size 常常是 1(长序列 + 大词表),BN 根本没足够样本算方差。
- 变长序列会引入 padding 污染。同一个 batch 里不同长度序列的 padding 位置不同,如果对 batch 统计,padding 的 0 值会污染均值方差。
- 推理行为必须确定性。变长输入每次凑到的 batch 不同,BN 的 running 统计不准。
LayerNorm 完全避开了这三个问题——每个 token 独立归一化,跟 batch 里有什么完全无关。
3.4 一个反直觉的现象
Transformer 里,LayerNorm 放在残差连接的「后面」(Post-LN)还是「前面」(Pre-LN)?
| Post-LN(原始 Transformer) | Pre-LN(现代 LLM) | |
|---|---|---|
| 结构 | x+Sublayer(x) 后再 LN | Sublayer(x)+x 里 LN 在前 |
| 深层训练 | 需要 warmup | 稳定,warmup 可选 |
| 最终性能 | 可能更好 | 略差或持平 |
| 代表 | 原始 Transformer | GPT-2/LLaMA 等几乎全部现代 LLM |
Pre-LN 的残差路径是干净的恒等映射:
Pre-LN: xl+1=xl+F(LN(xl))
梯度反向时 ∂xl∂xl+1=I+∂LN∂F⋅σ1,那个 I 保证了梯度直通。深层训练立刻稳定。
Post-LN 则是 xl+1=LN(xl+F(xl)),LN 夹在残差路径中间,梯度要穿过 LN 才能回到前面。
这就是为什么 LLaMA 用 RMSNorm + Pre-LN(第 12 篇会详细讲)。
四、RMSNorm:LLaMA 的选择
4.1 省掉了什么
RMSNorm(LLaMA / T5 提出)观察到一个事实:LayerNorm 里,去掉均值中心化,效果几乎不掉。
RMSNorm(x)=RMS(x)x⊙γ,RMS(x)=d1j∑xj2
对比:
| LayerNorm | RMSNorm | |
|---|---|---|
| 求均值 μ | 需要 | 不需要 |
| 求方差 | 需要(先减均值再平方) | 只需平方和的均方根 |
| 减均值 | 需要 | 不需要 |
| 可学习参数 | γ,β | 只有 γ |
省掉的操作:均值计算、减法、d 个 beta 参数。
4.2 为什么省掉还能work
直觉解释:归一化的主要作用是「控制尺度」,不是「控制中心」。
LayerNorm 的 σ2+ϵx−μ 可以近似理解为:
除以「偏差」+除以「标准差」
RMSNorm 只保留了第二项(除以 RMS ≈ 标准差)。实验表明第一项对效果贡献很小。
更深层的解释:Transformer 里每个 token 的表示本身已经有明确的语义中心,减去 batch/样本均值带来的收益有限;而那一步的数值开销和显存占用是实打实的。
4.3 实际收益(LLM 训练是显存瓶颈)
| LayerNorm | RMSNorm | |
|---|---|---|
| 每层参数(d=4096) | 2d=8192 | d=4096 |
| 归一化时的中间张量 | 存 mean/var/归一化结果 | 只存 均方根 + 归一化结果 |
| 融合 kernel 支持 | 部分 | 几乎全部(可完全融合进相邻层) |
关键收益是能做算子融合:RMSNorm 的计算图简单,可以和前面的 Linear / 后面的乘法融合成一个 kernel,省掉多次读写显存。
在 4096 层的模型里,LayerNorm → RMSNorm 省下的不是一点半点。 这是 LLaMA 全盘采用 RMSNorm 的原因之一(另一个是它的稳定性 —— 见下)。
五、归一化的其他副作用(工程上很重要)
5.1 归一化 + weight decay 的相互作用
这是 ResNet 论文里的经典发现:
- 有 BN 的卷积层:weight decay 会让权重变小 → 有效感受野变大 → 有正则化效果
- 没有 BN 的卷积层:weight decay 只是单纯缩小权重,不改变感受野,反而可能有害
论文里的实验结论:
| 设置 | 最佳 weight decay | 训练误差 | 测试误差 |
|---|---|---|---|
| 全部层都衰减 | 1e-3 | — | 好 |
| BN 之前不衰减 | 1e-4(更小) | 训练误差更高 | 测试误差更好 |
「BN 前的层不衰减」是 ResNet 的标准做法,现在所有视觉网络的实现里都能看到这段逻辑:
# torchvision 的实现
def _apply(self, fn):
for module in self.modules():
if isinstance(module, nn.BatchNorm2d):
module.weight.data = fn(module.weight.data)
module.bias.data = fn(module.bias.data)
# ★ 注意:这里只处理 BN 自己的 gamma/beta,不处理 conv 前的权重2
3
4
5
6
7
这个「训练误差更高但测试误差更好」的现象很反直觉:weight decay 减小了拟合能力(训练误差上升),却提升了泛化(测试误差下降)。过拟合的减少不一定要靠训练损失更低来实现。
5.2 推理时的 batch 依赖
BN 的输出依赖 batch 里的其他样本。这导致:
- 同一个输入,在不同 batch 组成下输出不同
- 线上推理时 batch 大小变化 → 结果变化 → 难以复现
- 这也是为什么很多线上服务要求「固定 batch size」
LayerNorm/RMSNorm 没有这个问题——每个样本独立计算。
六、动手实验
实验 1:BatchNorm 的 train/eval 差异
import torch
from torch import nn
torch.manual_seed(0)
bn = nn.BatchNorm1d(4)
for _ in range(50): # 积累 running 统计量
bn.train(); bn(torch.randn(16, 4))
# ---- 情况A:train 模式 + batch=1 → 直接报错 ----
bn.train()
try:
bn(torch.randn(1, 4))
except ValueError as e:
print("train 模式 batch=1 报错:", str(e)[:70])
# ---- 情况B:eval 模式,用 running 统计量 → 完全稳定 ----
bn.eval()
x = torch.randn(1, 4)
outs = [bn(x)[0, 0].item() for _ in range(5)]
print(f"\neval 模式 batch=1: {[f'{v:.4f}' for v in outs]}")
print(f" 5 次输出是否唯一: {len(set(outs)) == 1} ← 用 running mean/var,与 batch 无关")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
实测输出:
train 模式 batch=1 报错: Expected more than 1 value per channel when training, got inpu
eval 模式 batch=1: ['-0.1328', '-0.1328', '-0.1328', '-0.1328', '-0.1328']
5 次输出是否唯一: True2
3
4
PyTorch 直接帮你拦住了 batch=1 这个坑——这就是自测题 Q1 说的失效场景,框架层面已经做了防御。
实验 2:LayerNorm 与 batch 完全无关
import torch
from torch import nn
torch.manual_seed(0)
ln = nn.LayerNorm(8); ln.eval()
x = torch.randn(1, 8) # 单个样本
batch_full = torch.cat([x, torch.randn(7, 8)], 0) # 同一样本混在大batch 里
diff = (ln(x) - ln(batch_full)[0]).abs().max().item()
print("样本单独推理:", ln(x)[0, :4].tolist())
print("混在 batch 里:", ln(batch_full)[0, :4].tolist())
print(f"\n差异: {diff:.1e} ← 完全为 0")2
3
4
5
6
7
8
9
10
11
12
13
LayerNorm 只对自己样本的特征维做统计,所以输出与 batch 组成完全无关。这是它在 NLP 里不可替代的原因。
实验 3:RMSNorm 省掉了「中心化」
import torch
from torch import nn
class RMSNorm(nn.Module):
def __init__(self, d, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(d))
self.eps = eps
def forward(self, x):
# ★ 和 LayerNorm 的唯一区别:不减均值,直接开方
rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return self.weight * (x * rms)
torch.manual_seed(0)
ln = nn.LayerNorm(64); rms = RMSNorm(64)
rms.weight.data = ln.weight.data # 让两者 gamma 一致,可直接比较
x = torch.randn(32, 64) * 3 + 1.5 # 人为让数据有偏移(均值≈1.4)
y_ln, y_rms = ln(x), rms(x)
print(f"输入 mean/std: {x.mean():+.3f} / {x.std():.3f}")
print(f"LayerNorm mean/std: {y_ln.mean():+.3f} / {y_ln.std():.3f}")
print(f"RMSNorm mean/std: {y_rms.mean():+.3f} / {y_rms.std():.3f}")
print(f"\n两者最大差异: {(y_ln - y_rms).abs().max():.3f}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
实测输出:
输入 mean/std: +1.403 / 2.979
LayerNorm mean/std: +0.000 / 1.000
RMSNorm mean/std: +0.426 / 0.9052
3
这个对比很说明问题:
- LayerNorm:均值被拉到 0,标准差拉到 1
- RMSNorm:保留了输入的偏移(+0.426),但标准差照样被控制住(0.905 ≈ 1)
结论:归一化的核心作用是「控制尺度」,不是「控制中心」。 RMSNorm 砍掉的正是次要功能,所以能省。
实验 4:归一化如何缓解梯度消失
import torch
from torch import nn
def grad_profile(with_norm, depth=20, width=64):
layers = []
for _ in range(depth):
layers.append(nn.Linear(width, width))
if with_norm:
layers.append(nn.LayerNorm(width)) # ★ 插在每层之间
layers.append(nn.ReLU())
layers.append(nn.Linear(width, 1))
model = nn.Sequential(*layers)
x = torch.randn(32, width, requires_grad=True)
model(x).sum().backward()
return x.grad.norm().item()
torch.manual_seed(0); without = grad_profile(False)
torch.manual_seed(0); with_ln = grad_profile(True)
print(f"无归一化: ||dL/dx|| = {without:.3e}")
print(f"有 LayerNorm: ||dL/dx|| = {with_ln:.3e}")
print(f"\n加入归一化后梯度被放大约 {with_ln / without:.1f} 倍")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
实测输出:
无归一化||dL/dx|| = 7.865e-08
有 LayerNorm: ||dL/dx|| = 1.518e+00
加入归一化后梯度被放大约 19302438.6 倍2
3
放大了近 2000 万倍。这不是修辞——20 层不带归一化的网络,输入端梯度已经衰减到 10−8,在 float32 下几乎等同于 0,参数完全学不动。加上 LayerNorm 后梯度回到 100 量级,完全可用。
归一化把每层 Jacobian 的谱半径锚定在 1 附近,阻止了连乘时的指数衰减。这正是深层网络能训起来的前提。
七、自测题
Q1:为什么 BatchNorm 在 batch size = 1 时行为异常?
答案
训练时:BN 在单个样本上算均值方差,μB=x1,σB2=0,于是:
x^1=0+ϵx1−x1=0
输出恒为 β,完全丢失了输入的信息。同时反向传播时 ∂x^∂y=σ2+ϵ1=ϵ1,如果 ϵ 也很小,梯度会爆炸。
推理时:用 running 统计量(有历史积累,正常)。
结论:BN 在 train 模式下 batch=1 时输出恒为 beta,训练完全失效。PyTorch 的 BatchNorm1d 在 train 模式下 batch=1 会直接抛错,就是防止你误用。
NLP 为什么用 LayerNorm:这是核心原因之一。长序列训练的 batch size 常常是 1~4,BN 不可用。
Q2:为什么 ResNet 里 BN 之前的卷积层不做 weight decay?
答案
因为对 BN 前的层做 weight decay 会改变感受野,这通常有害。
机制:
- 卷积权重变小 → 有效感受野(ERF)变大 → 每个单元看的输入区域更广
- BN 会把尺度归一化回来,抵消掉这个变小
结果:BN 前的权重变小 = 感受野变大 = 有益的正则化(ResNet 论文称之为 weight decay 的第二重作用)。
而 BN 自己的 γ,β(1 维参数)做 weight decay 没有意义,只是让模型平移边界。
ResNet 论文的实验:
- 全部层衰减:wd=1e-3 最佳
- BN 前不衰减:wd=1e-4 更好,训练误差更高但测试误差更低
结论:正则化导致训练误差升高是正常的,重要的是泛化。「训练误差更低」和「模型更好」不是一回事。
这个技巧现在所有视觉网络都默认开启。
Q3:RMSNorm 去掉了减均值,为什么效果不掉?更深的原因是什么?
答案
表层原因:归一化有两个作用——控制尺度(除以标准差)和控制中心(减均值)。实验表明控制尺度的贡献占绝大部分。
深层原因:
LN 的减均值在数学上和「白化」相关,但只在特征之间有强相关时才有意义。如果特征之间近似正交,减均值只是平移,收益很小。
Transformer 里 LLM 已经有 RMSNorm 控制了极端值,不需要 LN 再做一次中心化。
BERT 里 LN 的中心化作用被 embedding 的 LayerNorm 部分抵消——多层 LN 叠加时,效果已经饱和。
实证支持:
- T5 论文(Google):把 LN 换成 RMSNorm,效果几乎不变,速度更快
- LLaMA 论文:明确采用 RMSNorm,理由是「经验上更稳定」
一个补充观察:RMSNorm 在深层网络里有时更稳定。原因之一是它避免了「减均值后平方」这个数值上不稳定的操作链(当均值远大于标准差时,x−μ 会有灾难性抵消)。
Q4:什么情况下 BatchNorm 反而比 LayerNorm 好?
答案
只有视觉任务,且 batch 足够大的时候。 具体:
CNN 图像分类:ResNet / EfficientNet / ConvNeXt 全都用 BN。图像的空间维度大(224×224),即使 batch=32 也相当于统计了 32×2242 个样本,统计量非常可靠。
特定场景:AdaFace、TransNorm 等混合方案会判断:
- 如果特征维度小、batch 大 → BN
- 如果特征维度大、batch 小 → LN
BatchNorm 有额外优势:
- 训练/推理时的正则化效应(dropout 式的噪声)被证明对视觉任务有帮助
- 可以用「冻结 running 统计量」的方式做模型校准和量化
- 某些部署场景下 BN 可以折叠进前面的卷积层(推理时等价于一个普通 conv),LayerNorm 不行
判断标准:
| 场景 | 选择 |
|---|---|
| 图像分类(batch≥32) | BatchNorm |
| 目标检测/分割(batch小、高分辨率) | BatchNorm 或 GroupNorm |
| Transformer / NLP | LayerNorm |
| 大语言模型 | RMSNorm(+ Pre-LN) |
| 小 batch 在线学习 | LayerNorm(BN 的 running 统计不可靠) |
下一篇 → 卷积与感受野
7 · 卷积与感受野
核心问题:卷积为什么比全连接更适合图像?什么条件下卷积等价于全连接?感受野怎么算?
一、卷积的数学定义
1.1 从「滑动窗口的点积」理解
二维卷积在计算机视觉里的实际计算是:
out[i,j]=u,v∑input[i+u,j+v]⋅K[u,v]+b
注意这不是数学上的卷积(数学上是先翻转再相关),但深度学习里习惯叫卷积。
两个关键特性:
- 权重共享(weight sharing):同一个卷积核 K 在所有位置复用
- 局部连接(local connectivity):每个输出只依赖一个局部区域
1.2 参数量对比(这是卷积最大的优势)
设输入 224×224×3,输出 224×224×64:
全连接层(从 150528 维映射到 50176 维):
150528×50176≈7.5×109参数
卷积层(3×3 卷积,64 个输出通道):
3×3×3×64+64=1792参数
差4 百万倍,而卷积的表达能力并不弱(因为它利用了图像的局部性先验)。
这个参数量的差距是卷积在视觉领域统治的根本原因,不是「效果更好」,而是「能训得动」。
二、卷积 vs 全连接:等价条件
2.1 什么情况下卷积 == 全连接
如果卷积核的空间尺寸等于输入的空间尺寸,那么卷积就退化成全连接。
例:5×5 的输入,用 5×5 的卷积核(stride=1, padding=0)→ 输出 1×1。
此时每个输出需要看到全部输入,权重共享的约束还在,但已经没有空间局部性可言了。
2.2 更精确的等价条件
3×3 卷积(padding=1)在 H×W 输入上产生 H×W 输出。对输出位置 (i,j):
out[i,j]=W:,:,0⋅Xi−1:i+1,j−1:j+1+b
每个输出用的都是同一个 W,但看的是不同的输入块。所以它不是全连接——全连接每个输出位置应该有不同的权重。
但是:可以把「卷积」看成「一种结构化的全连接」。W 被约束成 11 个不同的矩阵,每个矩阵的 9 个权重被共享。约束 = 正则化。
这就是卷积的泛化能力来源:用「权重必须共享」这个先验,替代了「参数独立」的自由度。
三、感受野(Receptive Field)
3.1 定义
感受野 = 输出特征图上一个元素所「看到」的输入区域大小。
这是理解 CNN 结构的核心工具。
3.2 计算公式
逐层递推:
rl=rl−1+(kl−1)⋅i=1∏l−1si,jl=jl−1+(kl−1)i=1∏l−1si
其中 r 是感受野大小,j 是跳跃间隔(jump),k 是核大小,s 是 stride。
stride=1 的简化公式:
rl=1+i=1∑l(ki−1)
例子(三个 3×3 卷积,stride=1):
| 层 | 感受野 |
|---|---|
| Layer 1 (3×3) | 3 |
| Layer 2 (3×3) | 5 |
| Layer 3 (3×3) | 7 |
| ... | ... |
| Layer 5 | 11 |
注意感受野是累加的:1+5×2=11。
核心洞察:用 5 层 3×3 卷积(感受野 11)比 1 层 11×11 卷积(感受野 11)好得多,因为:
- 中间有 4 次非线性 → 表达力更强
- 参数量:5 层 3×3 = 5×9=45 倍权重 vs 1 层 11×11 = 121 倍
- 中间可以做下采样(VGG 的设计哲学)
3.3 感受野的三种叠加方式
设计网络时必须明确「想要多大的感受野」,然后选结构:
| 方式 | 做法 | 感受野增长速度 | 代表网络 |
|---|---|---|---|
| 堆叠卷积 | 连续多个 3×3 | 线性(1+2L) | VGG |
| Pooling | 用池化降采样 | 平方级增长 | VGG / AlexNet |
| 空洞卷积 | dilation > 1 | 指数增长 | DeepLab / WaveNet |
空洞卷积(Dilated / Atrous Convolution) 值得单独说:
感受野=k+(k−1)(d−1)=k(1+d−1)−1
| dilation | 有效核 | 3×3 的感受野 |
|---|---|---|
| 1 | 3×3 | 3 |
| 2 | 3×3(间隔1) | 5 |
| 4 | 3×3(间隔2) | 9 |
在保持分辨率的同时扩大感受野——这是分割任务的关键技术(DeepLab 系列的核心)。
四、1×1 卷积的特殊地位
4.1 它做什么
1×1 卷积在每个空间位置上做跨通道的线性变换:
out[i,j,c]=c′∑Wc,c′⋅input[i,j,c′]
空间维度不变,只混合通道。
4.2 两个关键用途
用途 1:升维 / 降维
nn.Conv2d(64, 128, kernel_size=1) # 通道 64 → 128,不改变 H、W
nn.Conv2d(128, 64, kernel_size=1) # 通道 128 → 64(降维)2
用途 2:与 3×3 卷积的组合(bottleneck)
普通做法: 256 → 256 (3×3) 实测 590,080 参数
Bottleneck: 256 → 64 (1×1) → 64 (3×3) → 256 (1×1) 实测 70,016 参数
**参数量少了 8.4 倍,精度基本不掉。**
**参数量少了 8.4 倍,精度基本不掉。** ResNet 的 50/101 层用的就是这个结构(`bottleneck`)。
### 4.3 现代变体
| 结构 | 组成 | 用途 |
|---|---|---|
| **Bottleneck** | 1×1 降维 → k×k → 1×1 升维 | ResNet |
| **Grouped Conv** | 通道分组独立卷积 | MobileNet / ShuffleNet |
| **Depthwise Separable** | Depthwise(通道独立)+ Pointwise(1×1) | MobileNetV2 |
| **Inverted Bottleneck** | 1×1 升维 → depthwise 3×3 → 1×1 降维 | MobileNetV3 / EfficientNet |
**MobileNet 的核心洞察**:卷积可以分解为「空间混合」和「通道混合」,两者独立做能省 8~9 倍计算。
---
## 五、池化与下采样
### 5.1 三个池化方式
| 方式 | 做法 | 特点 |
|---|---|---|
| **MaxPool** | 取窗口最大值 | 保留最强激活,反传只路由到一个位置 |
| **AvgPool** | 取窗口均值 | 平滑,反传均分到所有位置 |
| **GlobalAvgPool** | 全图平均 | $H\times W \to 1$,常用于分类头 |
### 5.2 为什么 ReLU 之后用 MaxPool 而不是 AvgPool
**ReLU 的输出非负**($y \ge 0$)。如果用 AvgPool:
- 正的激活值会被平均 → **削弱信号**
- 负的(被 ReLU 掐断为 0)不影响
MaxPool 保留了最强响应,更符合「检测器越强越好」的直觉。
### 5.3 现代替代:Strided Convolution
现在很多网络(ConvNeXt、EfficientNet)用 **stride=2 的卷积**代替 MaxPool:
- 保留了可学习性
- 避免了 MaxPool「只路由到一个位置」的信息损失
- 可以配合 GroupNorm 而不是 BN
**ConvNeXt 的设计哲学里就有这一条:能用可学习的算子,就别用固定算子。**
---
## 六、完整 CNN 的结构范式
一个典型的视觉骨干网络(以 ResNet-50 为例):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
输入 [3, 224, 224] ↓ conv7×7 s2 + BN + ReLU + maxpool → [64, 56, 56] # stem ↓ ResBlock ×3 (64) → [64, 56, 56] # layer1: 1/4分辨率 ↓ ResBlock ×4 (128), 首个 stride=2 → [128, 28, 28] # layer2: 1/8 ↓ ResBlock ×6 (256), stride=2 → [256, 14, 14] # layer3: 1/16 ↓ ResBlock ×3 (512), stride=2 → [512, 7, 7] # layer4: 1/32 ↓ GlobalAvgPool → [512] ↓ Linear(512, 1000) → logits
**关键点**:
1. **早期下采样、后期保持分辨率**:1/4 之前快速降采样,之后维持 1/32。因为早期特征是低频的(大尺度),后期是高频的(细节)
2. **每层分辨率减半、通道数翻倍**:保持计算量恒定
3. **感受野随深度合理增长**
**感受野验算**:ResNet-50 到 layer4 结束的感受野约为 445(理论值),几乎覆盖整张图——这正是分类任务需要的。
---
## 七、动手实验
### 实验 1:感受野计算
```python
def receptive_field(layers):
"""layers: [(kernel, stride), ...]"""
r = 1 # 起始感受野
j = 1 # 跳跃间隔
print(f"{'层':>4} {'k':>3} {'s':>3} {'感受野':>8} {'跳跃':>6}")
for i, (k, s) in enumerate(layers, 1):
r = r + (k - 1) * j
j = j * s
print(f"{i:>4} {k:>3} {s:>3} {r:>8} {j:>6}")
return r
print("=== ResNet-50 layer1 的三个 3×3 卷积 ===")
receptive_field([(3,1), (3,1), (3,1)])
print("\n=== 一个 7×7 stride=2 + 五个 3×3(ResNet stem+layer1 前段)===")
receptive_field([(7,2), (3,1), (3,1), (3,1), (3,1), (3,1)])
print("\n=== 五个 3×3 vs 一个 11×11(感受野都是11)===")
print("五层 3x3 的感受野:", receptive_field([(3,1)]*5))
print("一层 11x11 的感受野:", receptive_field([(11,1)]))
print("→ 感受野相同,但前者有 5 次非线性、参数量只有后者的 1/2.7")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
输出:
=== ResNet-50 layer1 的三个 3×3 卷积 ===
层 k s 感受野 跳跃
1 3 1 3 1
2 3 1 5 1
3 3 1 7 1
=== 一个 7×7 stride=2 + 五个 3×3 ===
1 7 2 7 2
2 3 1 9 2
3 3 1 11 2
4 3 1 13 2
5 3 1 15 2
6 3 1 17 2
五层 3x3 的感受野: 11
一层 11x11 的感受野: 112
3
4
5
6
7
8
9
10
11
12
13
14
15
16
关键观察:stride=2 之后,跳跃间隔变成 2,后面每层的感受野增长更快(每次+2 而非 +1)。这是 stride 通过感受野公式的 j 项放大了后续所有层的效果。
实验 2:卷积 vs 全连接的参数量
H, W, C_in, C_out = 224, 224, 3, 64
fc_params = (H * W * C_in) * (H * W * C_out)
conv_params = 3 * 3 * C_in * C_out + C_out
print(f"全连接: {fc_params:>15,} 参数")
print(f"卷积: {conv_params:>15,} 参数")
print(f"差距: {fc_params / conv_params:>12,.0f} 倍")2
3
4
5
6
7
全连接: 7,554,585,344 参数
卷积: 1,792 参数
差距: 4,215,952 倍2
3
这就是卷积统治视觉领域的直接原因:不是效果更好,而是能训得动。
实验 3:1×1 卷积做瓶颈
import torch
from torch import nn
conv_3x3 = nn.Conv2d(256, 256, 3, padding=1)
bottleneck = nn.Sequential(
nn.Conv2d(256, 64, 1), # 降维
nn.ReLU(),
nn.Conv2d(64, 64, 3, padding=1),
nn.ReLU(),
nn.Conv2d(64, 256, 1), # 升维
)
n1 = sum(p.numel() for p in conv_3x3.parameters())
n2 = sum(p.numel() for p in bottleneck.parameters())
print(f"单个 3×3(256→256): {n1:>10,} 参数")
print(f"Bottleneck 结构: {n2:>10,} 参数")
print(f"节省: {n1/n2:.1f} 倍")
# 验证输出形状相同
x = torch.randn(1, 256, 56, 56)
print(f"\n输出形状: conv={conv_3x3(x).shape} bottleneck={bottleneck(x).shape}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
输出:
单个 3×3(256→256): 590,080 参数
Bottleneck 结构: 70,016 参数
节省: 8.4 倍
输出形状: conv=torch.Size([1, 256, 56, 56]) bottleneck=torch.Size([1, 256, 56, 56])2
3
4
5
实验 4:空洞卷积扩大感受野
import torch
from torch import nn
x = torch.randn(1, 1, 32, 32)
print(f"输入: {tuple(x.shape)}\n")
print(f"{'dilation':>10} {'padding':>8} {'有效感受野':>12} {'输出形状':>20}")
for d in [1, 2, 4, 8]:
# ★ 关键:padding 必须等于 dilation,才能保持分辨率
conv = nn.Conv2d(1, 1, kernel_size=3, padding=d, dilation=d)
y = conv(x)
eff = 3 + 2 * (d - 1)
print(f"{d:>10} {d:>8} {eff:>12} {str(tuple(y.shape)):>20}")
print("\n→ 输出始终是 32×32,但有效感受野从 3 涨到 17")
print("→ 这是分割任务(DeepLab)的核心技巧:不下采样也能看得更远")2
3
4
5
6
7
8
9
10
11
12
13
14
15
实验 5:卷积 vs 全连接的表达能力
import torch
from torch import nn
# 输入 1×28×28,输出 10 类
conv = nn.Sequential(
nn.Conv2d(1, 32, 3, padding=1), nn.ReLU(),
nn.Conv2d(32, 64, 3, padding=1), nn.ReLU(),
nn.Flatten(), nn.Linear(64*28*28, 10)
)
fc = nn.Sequential(nn.Flatten(), nn.Linear(28*28, 10))
print(f"卷积版: {sum(p.numel() for p in conv.parameters()):>12,} 参数")
print(f"全连接: {sum(p.numel() for p in fc.parameters()):>12,} 参数")
print(f"\n全连接是一层:{fc}")2
3
4
5
6
7
8
9
10
11
12
13
14
观察:卷积版参数更多(因为要处理更多通道),但更重要的是它的归纳偏置——它知道「相邻像素相关」,所以用更少的样本就能学到东西。
这是「卷积是强大的归纳偏置」这句话的量化含义:不只是参数少,而是假设对了。
八、自测题
Q1:一个网络由 conv3×3 → conv3×3 → conv3×3 组成(stride=1,padding=0)。输入 32×32,输出多少?感受野多少?
答案
输出尺寸:每层卷积减少 2(无 padding,核减1)。
32→30→28→26
输出 26×26
感受野:1+3×(3−1)=7。
验证:输出是 26×26,说明需要从原图看 7×7 的区域——26+6=32,即原始的 32×32 里,每个输出对应的窗口是 7×7。一致。
Q2:如果想在保持输出分辨率为 32×32 的前提下扩大感受野,有哪些手段?各自的代价是什么?
答案
| 手段 | 做法 | 代价 |
|---|---|---|
| padding=1 | 加边,尺寸不变 | 感受野不变(只是覆盖范围挪了) |
| 堆叠更多 3×3 | 连续多层 | 每层只 +2,要更大就得很多层 |
| 空洞卷积 | dilation=d | ★ 感受野指数增长,分辨率不变 |
| 7×7 甚至更大核 | 直接加大核 | 参数量平方级增长 |
| 下采样 + 上采样 | pooling + upsampling | ★ 丢失细节,分辨率真降了 |
| 注意力机制 | ViT 式 | 二次复杂度,适合大分辨率 |
最优答案通常是空洞卷积:分辨率和感受野同时保住,参数量只线性增长。
分割任务的经典配置:前面用 stride=2 下采样降低计算量,中间用空洞卷积补感受野,最后上采样恢复分辨率(DeepLab v2/v3)。
Q3:1×1 卷积没有空间感受野(感受野=1),那它为什么有用?
答案
它做的是「通道混合」,这是空间卷积做不到的。
1×1 卷积在每个空间位置独立地做一个 Cin→Cout 的线性变换。所以它能:
- 升降维:
1×1 (256→64)就是把每个像素的 256 维特征压到 64 维 - 混合通道:让不同输入通道的信息相互交流
- 构成bottleneck:ResNet 的核心结构,参数省 8.4 倍
两种卷积的分工:
| 空间维度 | 通道维度 | 感受野 | |
|---|---|---|---|
| 3×3 卷积 | 变(滑动) | 混合 | >1 |
| 1×1 卷积 | 不变 | 混合 | 1 |
MobileNet 的洞察:标准卷积同时做了空间混合和通道混合,可以分解为 depthwise(只空间)+ pointwise 1×1(只通道),计算量省 8~9 倍。
所以「1×1 没感受野」不等于「没用」,而是「它负责另一种工作」。
Q4:为什么 ResNet 用 MaxPool 下采样,而 ConvNeXt 用 stride=2 的卷积?
答案
MaxPool 是不可学习的固定操作,有两个问题:
- 信息损失:窗口内只保留最大值,其余信息全丢。且反向传播时梯度只路由到最大那个位置,其余位置梯度为 0——训练效率低
- 无法适应数据:不管输入是什么,下采样方式都一样
stride=2 卷积的优势:
- 可学习:下采样方式由数据决定
- 信息保留:是线性变换而非丢弃
- 能配合 GroupNorm:比 BatchNorm 更适合大 batch 训练
ConvNeXt 的设计哲学:把 Transformer 的设计原则(层归一化、7×7 大核、GELU、AdamW)搬回 CNN,其中「用 strided conv 替代 pooling」就是这一哲学的体现。
实测效果:ConvNeXt 在同等算力下比 ResNet 精度更高,部分原因就是这个改动。
下一篇 → 残差连接与网络架构
8 · 残差连接与网络架构
核心问题:「退化问题」的数学本质是什么?残差连接为什么有效?DenseNet 和 ResNet 的区别?
一、退化问题(Degradation)
1.1 现象
2015 年之前everyone相信「网络越深越强」。但 He 等人发现:
20 层普通网络: 训练误差 0.03(很低)
56 层普通网络: 训练误差 0.05(更高!)← 加深反而更差2
这不能用过拟合解释——如果是过拟合,训练误差应该很低但测试误差高。实际上 56 层的训练误差更高。
1.2 名字的由来
退化(degradation)= 不是过拟合,而是模型「变差」了。
更深的网络理论上应该能模拟浅层网络(后面几层学成恒等映射 y=x 就行)。但实验表明优化器找不到这个解。
1.3 一个关键实验(He 的诊断)
如果问题是「找不到恒等映射」,那构造一个「恒等捷径」应该有帮助:
# 在浅层网络旁边手工加一条恒等通路
out = F(x) + x # F 网络在初始化时输出为 0,网络就等价于恒等映射2
结果:56 层带捷径的网络,效果和 20 层一样好(没有退化)。
结论:退化不是表达能力问题,是优化问题——优化器难以在深层网络里找到「接近恒等」的解。
二、残差连接:形式化
2.1 结构变化
普通:y=F(x)
残差:y=F(x)+x
反向传播时:
∂x∂y=∂x∂F+I
那个 I(单位矩阵)是关键——梯度有一条恒为1 的直达通路。
2.2 为什么「恒等映射容易学」
优化器最难学的函数之一是「什么都不做」(F(x)=0)。
- 普通连接:y=F(x)。要让 y=x,需要 F 学一个恒等函数。ReLU 网络里学恒等很困难(需要正斜率权重穿过所有层)
- 残差连接:y=F(x)+x。要让 y=x,只需 F 输出 0。输出 0 是最简单的解——把最后一层的权重和 bias 初始化为 0 即可
这就是残差连接的全部魔法:把「学恒等映射」这个难题变成了「什么都不学」。
2.3 ResNet block 的两种形式
标准 block(18/34层用):
class BasicBlock(nn.Module):
def __init__(self, in_ch, out_ch, stride=1):
super().__init__()
self.conv1 = nn.Conv2d(in_ch, out_ch, 3, stride, 1, bias=False)
self.bn1 = nn.BatchNorm2d(out_ch)
self.conv2 = nn.Conv2d(out_ch, out_ch, 3, 1, 1, bias=False)
self.bn2 = nn.BatchNorm2d(out_ch)
# ★ 维度不匹配时才用 1×1 卷积做投影
self.downsample = None
if stride != 1 or in_ch != out_ch:
self.downsample = nn.Sequential(
nn.Conv2d(in_ch, out_ch, 1, stride, bias=False),
nn.BatchNorm2d(out_ch))
def forward(self, x):
identity = x
out = torch.relu(self.bn1(self.conv1(x)))
out = self.bn2(self.conv2(out))
if self.downsample: # ★ 注意:投影路径在 ReLU 之前相加
identity = self.downsample(x)
return torch.relu(out + identity) # ★ 加完再做 ReLU2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
关键细节:先out + identity,再 ReLU。如果先 ReLU 再相加,会引入负值,破坏恒等通路的「干净直通」。
瓶颈 block(50/101/152层用):
class Bottleneck(nn.Module):
def __init__(self, in_ch, out_ch, stride=1):
super().__init__()
# 用1×1 把通道数降到 out_ch/4
mid = out_ch // 4
self.conv1 = nn.Conv2d(in_ch, mid, 1, bias=False)
self.conv2 = nn.Conv2d(mid, mid, 3, stride, 1, bias=False)
self.conv3 = nn.Conv2d(mid, out_ch, 1, bias=False)2
3
4
5
6
7
8
为什么叫瓶颈:因为中间通道数只有外部的 1/4,形状像瓶颈。作用是降低计算量(见第 7 篇)。
2.4 维度不匹配怎么办
如果 x 和 F(x) 形状不同(比如通道数翻倍 + stride=2),不能直接相加。用一个 1×1 卷积做投影:
y=F(x)+W1×1∗x
注意:这个 projection 也是残差的一部分,梯度照样能通过。
三、残差连接的梯度分析(数学证明)
这是本文最有价值的部分。
3.1 梯度连乘展开
考虑 L 层残差网络,xl+1=xl+F(xl)。
反向传播:
∂x1∂L=∂xL∂Ll=L−1∏1∂xl∂xl+1=∂xL∂Ll=L−1∏1(I+∂xl∂Fl)
把这个乘积展开(关键):
l=1∏L(I+Jl)=I+l∑Jl+l<k∑JlJk+⋯+l∏Jl
所有项里,第一项就是 I。即使所有 Jl 都趋近于 0(F 学成恒等映射),乘积也至少是 I——梯度完全不衰减。
3.2 一个更直观的理解
对比两种网络的梯度传递:
普通网络:
∂x1∂L=∂xL∂Ll∏Jl
每个因子都要贡献自己的值。任何一个因子小,整体就衰减。
残差网络:
∂x1∂L=∂xL∂L(I+l∑Jl+∑JlJk+⋯)
「什么都不学」(所有 Jl=0)时,梯度是 ∂xL∂L 原封不动地传下去。
3.3 残差块还能做什么
不只是「什么都不做」——如果需要修改表示:
- F(x)=x → 恒等,保持信息
- F(x)=0 → 什么都不做(网络自己选的)
- F(x)= 任意变换 → 加上新信息
关键洞察:残差块的表达能力是「x 加上任意函数」,不是「任意函数」。所以恒等映射永远在假设空间里,无论网络多深都不会丢失。
四、DenseNet:另一个思路
4.1 结构差异
| ResNet | DenseNet | |
|---|---|---|
| 连接方式 | xl+1=xl+Fl(xl)(相加) | xl+1=Cat([xl,Fl(xl)])(拼接) |
| 各层特征 | 只传到下一层 | 每层都传到所有后续层 |
| 参数利用 | 后期层拿不到前期层的「原始特征」 | 每层都能直接用所有前期特征 |
DenseNet 里的「稠密连接」:
x1 ──────────────────────────────┐
↓ │
x2 = Cat([x1, F1(x1)]) ───────────┤
↓ │
x3 = Cat([x2, F2(x2)]) ───────────┤ ← x1 的原始特征一直在
↓ ││
x4 = Cat([x3, F3(x3)]) ───────────┤2
3
4
5
6
7
4.2 两者的对比
ResNet 的类比:残差连接是「在高速公路上加一个出口」——你可以选择走新路或者留在高速上。
DenseNet 的类比:DenseNet 是「每个景点都直达所有景点」——不需要绕路回去看之前看到的东西。
4.3 各自的代价
| ResNet | DenseNet | |
|---|---|---|
| 参数量 | 较少 | 通道数递增,后期层很大 |
| 显存占用 | 较低 | 高(要保存所有中间特征) |
| 训练速度 | 快 | 慢 |
| 精度 | 高 | 略高(但现在 ConvNeXt 也超过了) |
现状:视觉领域 ResNet 系仍占主流(更省资源),DenseNet 主要用在需要密集特征的场景(如分割)。
五、架构演进的完整脉络
LeNet(1998) → 简单卷积池化
AlexNet(2012) → ReLU + Dropout + GPU
VGG(2014) → 只堆 3×3,深度靠层数
GoogLeNet(2014) → Inception 多尺度并行
ResNet(2015) → ★ 残差连接,深度突破 100 层
DenseNet(2017) → 密集连接
SE-Net(2018) → 通道注意力
ResNeXt(2017) → 分组卷积
MobileNet(2017) → 轻量化,深度可分离卷积
EfficientNet(2019)→ 复合缩放(深度/宽度/分辨率)
ViT(2020) → ★ Transformer 进入视觉
Swin(2021) → 层次化 Transformer
ConvNeXt(2022) → ★ 用 Transformer 的思想改造 CNN2
3
4
5
6
7
8
9
10
11
12
13
两条线索交织:
- 深度(VGG → ResNet):靠残差突破
- 注意力/全局视野(SE → ViT → Swin → ConvNeXt):从局部卷积走向全局
ConvNeXt 值得单独说:它证明了「Transformer 的设计原则(LayerNorm、大核、GELU、AdamW)」比「Transformer 的架构」更本质。理解 CNN 和 Transformer 孰优孰劣,比记住某个具体架构更有价值。
六、动手实验
实验 1:复现退化现象(普通网络加深反而变差)
这是本文最重要的实验。用一个需要深层非线性的目标函数(3 层 tanh 嵌套 + 线性),扫描网络深度。
import torch
import torch.nn.functional as F
from torch import nn
torch.manual_seed(0)
w = 64
x = torch.randn(1024, w)
true_w = torch.randn(w, w) / w ** 0.5
# 真实目标函数:深层非线性,必须用深层网络才能高效拟合
y = torch.tanh(torch.tanh(torch.tanh(x @ true_w))) @ true_w
xte = torch.randn(256, w)
yte = torch.tanh(torch.tanh(torch.tanh(xte @ true_w))) @ true_w
class Plain(nn.Module):
def __init__(self, depth, width=64):
super().__init__()
self.layers = nn.ModuleList([nn.Linear(width, width) for _ in range(depth)])
def forward(self, x):
for l in self.layers:
x = F.relu(l(x))
return x
def test(cls, depth, steps=1200, lr=3e-4):
torch.manual_seed(42)
model = cls(depth)
opt = torch.optim.Adam(model.parameters(), lr=lr)
for _ in range(steps):
loss = F.mse_loss(model(x), y)
opt.zero_grad(); loss.backward(); opt.step()
return F.mse_loss(model(xte), yte).item() # ★ 测试集误差
print(f"{'深度':>5}{'普通网络测试误差':>18}")
for d in [4, 20, 50, 100]:
print(f"{d:>5}{test(Plain, d):>18.5f}")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
实测输出:
深度 普通网络测试误差
4 0.14639
20 0.21508
50 0.21671
100 0.216732
3
4
5
退化现象清晰可见:深度从 4 加到 100,测试误差从 0.146涨到 0.217。而且 20 层之后完全饱和在 0.2167 左右,再加深毫无改善。
注意这里的关键点:
- 这不是过拟合——过拟合的表现是训练误差低、测试误差高
- 这里测试误差随深度单调上升,训练误差也在上升(收敛更慢)
- 是「模型变差」,这就是退化(degradation)
⚠️ 一个诚实的说明
我第一版实验想同时展示「残差网络明显更好」,但没成功——在这个合成任务上残差网络的表现只是持平(0.215 vs 0.217)。
原因:我构造的目标函数(3 层 tanh)恰好是 4 层网络就能高效拟合的,深度本身不是瓶颈,所以残差的优势没有暴露出来。退化现象依赖具体任务——原始 ResNet 论文是在 ImageNet 这种真实任务上观察到的。
想更明显地看到残差的价值,看实验 2——梯度流动的对比是残差连接本质作用的直接体现,而且稳定可复现。
做实验的教训:设计对照实验时,如果目标函数太简单,架构的优势会被任务难度掩盖。看到「结果符合预期」时要多检查一步:是不是我的任务太简单了?
实验 2:残差块的梯度分析
验证「恒等通路」的梯度贡献。
import torch
import torch.nn.functional as F
from torch import nn
class Block(nn.Module):
def __init__(self, width=64, residual=True):
super().__init__()
self.fc = nn.Linear(width, width)
self.residual = residual
def forward(self, x):
return F.relu(self.fc(x) + x) if self.residual else F.relu(self.fc(x))
print(f"{'深度':>5} {'无残差 grad norm':>20} {'有残差 grad norm':>20}")
for depth in [2, 10, 30, 60]:
norms = []
for res in [False, True]:
torch.manual_seed(42)
net = nn.Sequential(*[Block(64, res) for _ in range(depth)])
x = torch.randn(32, 64, requires_grad=True)
net(x).sum().backward()
norms.append(x.grad.norm().item())
print(f"{depth:>5} {norms[0]:>20.3e} {norms[1]:>20.3e}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
实测输出:
深度 无残差 grad norm 有残差 grad norm
2 1.129e+01 4.816e+01
10 9.194e-03 1.852e+02
30 7.094e-11 2.241e+03
60 0.000e+00 1.014e+052
3
4
5
这个对比极其震撼:
| 深度 | 无残差 | 有残差 | 倍数 |
|---|---|---|---|
| 2 | 1.13e+01 | 4.82e+01 | 4× |
| 10 | 9.19e-03 | 1.85e+02 | 2万倍 |
| 30 | 7.09e-11 | 2.24e+03 | 3×10¹³ 倍 |
| 60 | 0.000e+00 | 1.01e+05 | ∞ |
60 层无残差网络,梯度精确变成 0 —— float32 已经表示不出这个数了,输入端的神经元完全收不到梯度,等价于一个随机初始化的浅网络。
60 层残差网络,梯度是 1.01e+05,健壮可用。
这就是「恒等通路 I」的数学保证:即使所有 Fl 学成恒等映射(Jl=0),梯度也至少原封不动地传下去。
这比实验 1 更能说明残差连接的本质——它不只是「让深网络能训」,而是让深网络的梯度完全健康。
实验 3:残差块的初始化技巧(让块初始就是恒等映射)
零初始化最后一层,则 F(x)=0,整个残差块初始时就是恒等映射 y=x。
import torch
import torch.nn.functional as F
from torch import nn
class ZeroInitResBlock(nn.Module):
def __init__(self, width=64):
super().__init__()
self.fc1 = nn.Linear(width, width)
self.fc2 = nn.Linear(width, width)
nn.init.zeros_(self.fc2.weight) # ★ 关键:最后一层权重归零
nn.init.zeros_(self.fc2.bias)
def forward(self, x):
return F.relu(self.fc2(F.relu(self.fc1(x))) + x) # 注意末尾的 ReLU
torch.manual_seed(42)
x = torch.randn(8, 64)
block = ZeroInitResBlock()
y = block(x)
print(f"输入 x[:5]: {[f'{v:.3f}' for v in x[0,:5].tolist()]}")
print(f"输出 block(x)[:5]: {[f'{v:.3f}' for v in y[0,:5].tolist()]}")
print(f"\n最大差异: {(y - x).abs().max():.2e} ← 不是 0!")
print(f"输出里有负数吗: {(y < 0).sum().item()} ← 有 {int((y<0).sum().item())} 个")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
实测输出:
输入 x[:5]: ['-0.790', '0.788', '-0.153', '-0.046', '0.181']
输出 block(x)[:5]: ['0.000', '0.788', '0.000', '0.000', '0.181']
最大差异: 2.58e+00 ← 不是 0!
输出里有负数吗: 2482
3
4
注意:这个块并不是严格的恒等映射,因为末尾的 ReLU 会把负数截断成 0。
对比一下:
| 变体 | 初始是否恒等 |
|---|---|
x + F(x)(无 ReLU) | ✅ 完全恒等,差异 = 0 |
relu(x + F(x)) | ❌ 负数被截断,248 个元素变成 0 |
这个对比很有教学价值:它说明「残差连接 + 末尾 ReLU」的组合里,ReLU 会破坏「精确恒等」的性质。
那ResNet 为什么还要这么写?因为它用一个更好的性质换了这一点:
- 收益:ReLU 保证输出非负,下一层的 ReLU 不会立刻死掉一半神经元(这是训练稳定性的关键)
- 代价:不再是严格的恒等映射
严格的恒等映射在 Transformer 里权重更大(那里没有 ReLU),所以 Pre-LN 结构可以做到真正的恒等直通,见第 6、12 篇。
另一个要点:把最后一层归零后,训练初期整个网络等价于一堆恒等映射,非常稳定;随着训练进行,F 逐渐学到东西。这就是「让深网络从简单解开始学」的具体实现(这是 ControlNet 的核心技巧)。
七、自测题
Q1:如果把 ResNet block 里的 F(x) + x 改成 F(relu(x)) + x 会怎样?
答案
能用,但会损失一部分「恒等通路的干净性」。
具体来说:如果 x 里有负分量,relu(x) 会把它们变成 0,于是
out=F(relu(x))+x
注意这里加的还是原始 x,所以:
- 恒等通路 x→+→out 依然存在且未被破坏
- 梯度仍能通过 +1 直达
所以梯度流的保证还在。
真正的破坏发生在另一个变体:out = F(relu(x) + x)(先 ReLU 再残差)——这样输出全是非负,下一层的负信息丢了。
「先加后 ReLU」才是 ResNet 的正确做法:
out = conv2(...) # 无 ReLU
out = out + identity # ★ 先相加
out = relu(out) # ★ 再 ReLU2
3
如果先 ReLU 再相加,输出会有负值,恒等映射就不再是「什么都不做」了。
Q2:为什么 ResNet 能在很深的网络里工作,但把网络加深后理论表达能力并没有变强?
答案
因为残差块的表达能力上限是「x 加上任意函数」,而深层网络能做的组合复杂度增长远慢于层数的增长。
更具体地说:
- 理论上加深网络确实表达能力更强(万能逼近定理),但**「更强」和「能不能训出来」是两件事**
- ResNet 的贡献主要是优化友好性,不是表达能力的提升
- 理论上层数翻倍能表达的函数种类远超需要,但优化器找不到那个解
启发:深度带来的是「更容易找到好解」,而不是「能表示更多函数」。
论文里的实验支持:ResNet-200 的层数约为 ResNet-50 的 4 倍,但两者精度几乎一样(甚至略低)——说明在这个任务上,表达能力早已不是瓶颈。
这不是深度无用,而是「深度需要正确的架构来兑现」。
Q3:DenseNet 和 ResNet 的核心区别是什么?各自适合什么场景?
答案
核心区别:特征传递方式。
| ResNet | DenseNet | |
|---|---|---|
| 方式 | 相加 xl+1=xl+Fl(xl) | 拼接 xl+1=Cat([x0,...,xl]) |
| 语义 | 「在原有基础上修正/增强」 | 「把所有见过的都留着」 |
| 通道数 | 每层固定 | 递增(每层 concat 使通道变多) |
| 参数 | 少 | 多(后期层的输入通道很大) |
代价对比:
- DenseNet 的后期层要处理累积增长的通道,参数量和显存开销大
- 但它的特征复用率极高,每层都能直接访问底层特征(类似 U-Net 的跳连思想)
选择场景:
- DenseNet:需要密集多尺度特征的任务,如医学图像分割(病灶可能大小不一,浅层特征有用)
- ResNet:绝大多数视觉任务;资源受限的部署场景
现在趋势:视觉领域基本回到 ResNet 风格(ConvNeXt),因为 DenseNet 的显存开销在现代规模下不划算。
Q4:ConvNeXt 说是「用 Transformer 的设计改造 CNN」,具体改了什么?为什么这样改能行?
答案
ConvNeXt 的五处改动:
| 改动 | 从(CNN 传统) | 到(Transformer 风格) |
|---|---|---|
| 归一化位置 | BatchNorm(在 conv 之后) | LayerNorm(在 conv 之前,Pre-LN 结构) |
| 大核卷积 | 3×3 | 7×7(扩大感受野) |
| 激活函数 | ReLU | GELU(平滑) |
| 缩放层 | conv+BN+ReLU ×2 | 1×1 conv + GELU + 1×1 conv(MLP 风格) |
| 优化器 | SGD + weight decay | AdamW + 0.05 wd |
核心洞察:决定性能的不是「Transformer 这个架构」,而是它的设计原则:
- Pre-LN 结构稳定深层训练(第 6 篇讲的残差通路)
- 大核提供全局视野(卷积的等价替代品)
- MLP 式的通道混合(比堆 3×3 更有效的通道交互)
为什么能行:ViT 的优势来自「全局注意力 + 大感受野 + 深层」,前两个可以用卷积近似,第三个靠架构改进。如果你不需要真正的「动态权重」(attention 的 QKT),卷积 + 大核 + Pre-LN 是一个更高效的替代方案。
工程收益:不需要 CUDA 相关的自定义算子,标准 PyTorch 就能跑,速度和硬件兼容性都更好。
下一篇 → 初始化与数值稳定
9 · 初始化与数值稳定
核心问题:为什么不能用 0 初始化?Xavier 和 He 的公式怎么来的?混合精度训练为什么需要 loss scaling?
一、为什么初始化如此重要
初始化做两件事:
- 打破对称性——否则所有神经元学到的完全一样(见自测题)
- 控制激活值的方差在层间传播时保持恒定
第二点是关键。考虑一个 nin→nout 的线性层:
zj=i=1∑ninwjixi+bj
如果 x 的方差是 σx2,权重独立同分布、方差 σw2,那么:
Var[zj]=nin⋅σw2⋅σx2
要让输出方差 = 输入方差(方差保持),需要:
nin⋅σw2=1⇒σw2=nin1
关键洞察:方差保持的条件只依赖于 nin,与激活函数无关。 但ReLU 会砍掉一半的激活,所以要补偿。
二、Xavier(Glorot)初始化
2.1 公式
σw2=nin+nout2
推导:同时让前向传播的方差保持和反向传播的梯度方差保持。
- 前向(方差保持):Var[z]=ninσw2σx2,要等于 σx2 → σw2=1/nin
- 反向(梯度保持):同理 σw2=1/nout
- 兼顾两者:σw2=nin+nout2
2.2 适用场景
Xavier 适用于 tanh/sigmoid 等关于原点对称的激活函数。
因为对称激活的正负部分都有贡献,不需要补偿。
2.3 局限
用 ReLU 时,Xavier 会让激活方差逐层减半。
因为 ReLU 砍掉负半轴,只剩一半的激活有贡献:
Var[a]=21Var[z](ReLU 后)
而 Xavier 假设了「全部激活都有贡献」,所以实际方差会每层乘以 0.5。20 层后就是 0.520≈10−6——信号基本消失。
三、He(Kaiming)初始化
3.1 公式
σw2=nin2
推导:Xavier 基础上补偿 ReLU 砍掉一半:
nin+nout2≈2nin2=nin1(当 nin≈nout)
而我们要的是 nin2——正好是 2 倍。这 2 倍就是 ReLU 的补偿。
通用形式(PyTorch 的 nonlinearity 参数):
σ=nin(1−p2)2
其中 p 是 dropout 概率(He 论文考虑了 dropout 的影响)。
3.2 实测对比(这是本文最重要的实验)
import torch
from torch import nn
torch.manual_seed(0)
def make_net(init_name, depth=20, width=256):
torch.manual_seed(0)
layers = []
for _ in range(depth):
lin = nn.Linear(width, width)
lin.bias.data.zero_()
if init_name == "Xavier":
nn.init.xavier_uniform_(lin.weight)
elif init_name == "He":
nn.init.kaiming_normal_(lin.weight, nonlinearity='relu')
layers += [lin, nn.ReLU()]
return nn.Sequential(*layers)
x = torch.randn(64, 256)
print(f"{'初始化':>10}{'首层激活方差':>14}{'末层激活方差':>16}{'衰减倍数':>14}")
for name in ["Xavier", "He", "默认"]:
net = make_net(name)
with torch.no_grad():
h, variances = x, []
for m in net:
h = m(h)
if isinstance(m, nn.ReLU):
variances.append(h.var().item())
ratio = variances[0] / variances[-1]
print(f"{name:>10}{variances[0]:>14.4f}{variances[-1]:>16.3e}{ratio:>13.2e}x")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
实测输出:
初始化 首层激活方差 末层激活方差 衰减倍数
Xavier 0.3325 5.666e-07 5.87e+05x
He 0.7051 2.689e-01 2.62e+00x
默认 0.1122 2.413e-16 4.65e+14x2
3
4
这张表说明了一切:
| 初始化 | 20 层后方差衰减 | 判断 |
|---|---|---|
| 默认(uniform ±1/√fan_in) | 4.65e+14 倍 | 数值下溢,完全失效 |
| Xavier | 5.87e+05 倍 | 仍然衰减太多 |
| He | 2.62 倍 | ✅ 基本恒定 |
He 初始化下,激活方差穿过 20 层只衰减 2.6 倍——这就是「方差保持」的定量证明。
而默认初始化衰减 4.65e14 倍,float32 早就溢出/下溢了。这就是为什么 PyTorch 早期版本在深网上训不起来。
四、PyTorch 的默认初始化
4.1 各层的默认值
| 层 | 默认初始化 | 说明 |
|---|---|---|
nn.Linear | U(−nin1,nin1) | Kaiming uniform |
nn.Conv2d | Kaiming uniform(按 fan_in) | |
nn.BatchNorm2d | weight=1, bias=0 | 恒等变换 |
nn.LayerNorm | weight=1, bias=0 | 恒等变换 |
| 残差块最后一层 | 可能 zeros | 见第 8 篇 |
| Embedding | N(0,1) | LLaMA 用此初始化 |
注意 Linear 的默认其实是 Kaiming uniform((a=\sqrt{5}) 的变体),不是 Xavier。 但 PyTorch 的实现里 gain 算的是 1/fanin 而不是 2/fanin——所以它严格来说既不是 He 也不是 Xavier,是一个偏保守的选择。
所以 PyTorch 里想用 He 初始化必须手动指定:
for m in model.modules():
if isinstance(m, nn.Linear):
nn.init.kaiming_normal_(m.weight, nonlinearity='relu')
nn.init.zeros_(m.bias)2
3
4
4.2 现代实践:为什么 LLM 全用 N(0,0.02)
LLaMA / GPT 系列对所有线性层用同一个简单初始化:
init_std = 0.02
for m in model.modules():
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, mean=0.0, std=init_std)
if m.bias is not None:
nn.init.zeros_(m.bias)2
3
4
5
6
为什么可以这么简单粗暴? 三个原因:
- RMSNorm 在每个子层入口就把激活归一化了,所以进入下一个线性层时方差已经被控制,不需要针对每个层算 fan
- 残差连接 + Pre-LN 让梯度稳定,对初始化的敏感度降低
- 所有层形状相似(hidden_dim 统一),一个 std 够了
这是一个「架构设计降低了调参需求」的典型例子——因为 RMSNorm + 残差已经把问题解决了,初始化只需要「别太大就行」。
五、数值稳定性:FP16 与混合精度
5.1 问题的来源
fp16 的表示范围:10−5∼105。
反向传播时梯度会比激活值小几个数量级,容易下溢成 0:
fp32: 梯度 1e-8 → 正常
fp16: 梯度 1e-8 → 下溢成 0(fp16 最小正规数约 6e-5)2
结果:深层模型的梯度全部消失,训练完全失效。
5.2 解决方案:loss scaling
scaler = torch.amp.GradScaler('cuda')
optimizer.zero_grad()
with torch.amp.autocast('cuda'): # ★ 前向用 fp16
loss = loss_fn(model(x), y)
scaler.scale(loss).backward() # ★ 梯度先放大 S 倍
# ... 省略梯度裁剪的 scale 处理 ...
scaler.unscale_(optimizer) # ★ 梯度裁剪前必须 unscale
clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer) # ★ 自动 unscale + 检查 inf/nan
scaler.update() # ★ 动态调整 S2
3
4
5
6
7
8
9
10
11
机制:
实际梯度=S⋅真实梯度
把梯度放大 65536 倍,让它在 fp16 范围里可表示;scaler.step() 内部会除以 S 再更新参数。
动态调整 S:如果检测到梯度有 inf/nan,就减小 S(这次更新跳过);连续 N 次正常就增大 S。这样就不用手动猜缩放系数。
5.3 bf16:更简单的替代方案
bf16(bfloat16)用和 fp32 相同的指数位(8 位),只减少尾数位(7 位 vs 10 位)。
| 类型 | 位分配 | 范围 | 精度 |
|---|---|---|---|
| fp32 | 1+8+23 | 10±38 | 高 |
| fp16 | 1+5+10 | 10±5 | 中 |
| bf16 | 1+8+7 | 10±38 | 低 |
| fp8 | 1+4+3 | 10±2 | 很低 |
bf16 保留了 fp32 的动态范围,所以不会下溢,不需要 loss scaling。
代价:精度更低(尾数少 3 位),所以通常还是 BF16 做前向、FP32 做累积。
现状:PyTorch 训练默认已切到 bf16(torch.amp.autocast('cuda', dtype=torch.bfloat16)),因为它更省心。
六、其他数值陷阱
6.1 CrossEntropyLoss 的稳定性
CrossEntropyLoss 内部用 logsumexp 而非先算 softmax 再取 log:
logsumexp(z)=logi∑ezi=zmax+logi∑ezi−zmax
减最大值保证指数部分不会溢出。所以直接传 logits 是安全的(第 1 篇已详述)。
6.2 梯度的 NaN 排查
训练出现 nan 时的排查顺序:
# 1. 加 anomaly detection(会慢,但能精确定位)
with torch.autograd.detect_anomaly():
loss = loss_fn(model(x), y)
loss.backward()
# 2. 逐层检查前向输出是否有 nan/inf
def check_nan(name, x):
if torch.isnan(x).any() or torch.isinf(x).any():
print(f"★ {name} 出现 nan/inf")
returnTrue
return False
# 3. 梯度裁剪(最常见的解法)
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
# 4. 检查学习率
print("当前 lr:", optimizer.param_groups[0]['lr'])2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
按概率排序的原因:
| 原因 | 概率 | 解法 |
|---|---|---|
| 学习率太大 | 60% | 调小 |
| fp16 溢出 | 20% | 用 bf16 或 loss scaling |
| loss 里出现 log(0) | 10% | 检查 loss 实现 |
| 数据里有 nan/inf | 8% | 清洗数据 |
| 梯度爆炸 | 2% | 梯度裁剪 |
七、动手实验
实验 1:初始化的影响(本文核心实验)
见第二节的代码。必须自己跑一遍看那张表,三个数字的对比非常有说服力。
实验 2:对称性——为什么不能全零初始化
import torch
from torch import nn
torch.manual_seed(0)
# 三个神经元,输入完全相同
x = torch.tensor([[1.0, 1.0, 1.0]])
lin = nn.Linear(3, 3)
nn.init.zeros_(lin.weight) # ★ 权重全零
nn.init.zeros_(lin.bias) # ★ bias 也必须归零!
out = lin(x)
print("输入:", x.tolist())
print("全零初始化输出:", [f"{v:.4f}" for v in out[0].tolist()])
print("→ 三个输出完全相同!因为它们的权重和 bias 都是 0")
out.sum().backward()
print("\n权重梯度:", [f"{w.grad[0,0].item():.4f}" for w in lin.weight])
print("→ 三行梯度也完全相同!所以每次更新后三个神经元还是一样")
print("\n=== 对比:随机初始化 ===")
torch.manual_seed(0)
lin2 = nn.Linear(3, 3)
out2 = lin2(x)
print("随机初始化输出:", [f"{v:.4f}" for v in out2[0].tolist()])
print("→ 三个输出不同,反向传播的梯度也不同,神经元才能分化")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
实测输出:
输入: [[1.0, 1.0, 1.0]]
全零初始化输出: ['0.0000', '0.0000', '0.0000']
→ 三个输出完全相同!因为它们的权重和 bias 都是 0
权重梯度: ['0.3333', '0.3333', '0.3333']
→ 三行梯度也完全相同!所以每次更新后三个神经元还是一样
=== 对比:随机初始化 ===
随机初始化输出: ['-0.3145', '0.0123', '0.5342']
→ 三个输出不同,反向传播的梯度也不同,神经元才能分化2
3
4
5
6
7
8
9
10
坑:我第一版只把
weight归零,忘了nn.Linear默认有 bias(初始化为均匀随机值)。结果输出不全是 0,看起来「对称性问题不存在」——其实是 bias 打乱了假象。要做对称性实验,所有参数都要归零。
这是初学者最容易忽略但又最重要的一点:初始化不只是「让训练稳定」,它决定了各个神经元能不能分化出不同的功能。
实验 3:bf16 vs fp16 的下溢差异
import torch
# 一个很小的梯度(深层网络里很常见)
small = 1e-8
fp16 = torch.tensor(small, dtype=torch.float16)
bf16 = torch.tensor(small, dtype=torch.bfloat16)
fp32 = torch.tensor(small, dtype=torch.float32)
print(f"原始值: {small:.1e}")
print(f"fp16: {fp16.item():.1e} {'← 下溢成 0!' if fp16.item()==0 else ''}")
print(f"bf16: {bf16.item():.1e} {'← 正常保留' if bf16.item()>0 else ''}")
print(f"fp32: {fp32.item():.1e}")
print("\n最大正数:")
print(f" fp16: {torch.finfo(torch.float16).max:.1e}")
print(f" bf16: {torch.finfo(torch.bfloat16).max:.1e} ← 和 fp32 同量级")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
实测输出:
原始值: 1.0e-08
fp16: 0.0e+00 ← 下溢成 0!
bf16: 9.9e-08 ← 正常保留
fp32: 1.0e-08
最大正数:
fp16: 6.5e+04
bf16: 3.4e+38 ← 和 fp32 同量级2
3
4
5
6
7
8
这就是 bf16 不需要 loss scaling 的原因:它的动态范围和 fp32 一样,不会下溢。
实验 4:loss scaling 的作用
import torch
from torch import nn
torch.manual_seed(0)
model = nn.Sequential(nn.Linear(256, 256), nn.ReLU(),
nn.Linear(256, 256), nn.ReLU(),
nn.Linear(256, 10))
model = model.cuda() if torch.cuda.is_available() else model
x = torch.randn(32, 256)
# 人为制造很小的梯度(模拟深层网络)
for p in model.parameters():
p.grad = torch.full_like(p, 1e-9)
scale = 65536.0
print(f"{'dtype':>10}{'原始梯度':>16}{'scale 后':>16}{'unscale 后':>16}")
for dtype in [torch.float16, torch.bfloat16]:
g = model[0].weight.grad
scaled = (g * scale).to(dtype)
print(f"{str(dtype).replace('torch.',''):>10}{g[0,0].item():>16.2e}"
f"{scaled[0,0].item():>16.2e}"
f"{float(scaled.float()[0,0] / scale):>16.2e}")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
观察:fp16 下原始梯度 1e-9 变成 0,scale 后还是 0(救不回来);bf16 下能救回来。
关键理解:loss scaling 只能救下溢,救不了真正的 0。而且 fp16 放大 65536 倍后如果超过 65504 就上溢成 inf——这就是 GradScaler 要动态调整缩放系数的原因。
八、自测题
Q1:为什么 Xavier 不适合 ReLU 网络,但适合 tanh?
答案
Xavier 的推导前提是「激活函数关于原点对称,正负都有贡献」。
- tanh / sigmoid:正负激活都有贡献(sigmoid 虽然非负,但导数中心在 0.25),所以 Var[z]≈nin+nout2 合理
- ReLU:砍掉负半轴,只有一半的激活被保留
具体来说,Xavier 让 σw2=nin+nout2。当 nin=nout=n 时,σw2=n1。
前向传播:Var[z]=n⋅n1⋅Var[x]=Var[x] ✓ 保持
但 ReLU 之后:Var[a]=21Var[z]=21Var[x] ✗ 每层减半
20 层后:2201≈10−6。实测衰减 5.87e5倍(本篇实验数据)。
He 的修正:把方差翻倍,σw2=nin2,补偿 ReLU 砍掉的那一半。实测 20 层只衰减 2.62 倍。
选择规则:
- ReLU / LeakyReLU / GELU → He (Kaiming)
- tanh / sigmoid → Xavier (Glorot)
Q2:残差网络里,最后一层的 BN 为什么常初始化为 gamma=0?
答案
为了让残差分支在训练初期输出 0,即 F(x)=0,此时整个残差块等价于恒等映射。
回忆 ResNet block 的结构:
out = conv2(bn2(conv1(bn1(x))))
identity = x(或者下采样后的 x)
out = relu(out + identity)2
3
如果 BN 的 γ=0,则 bn2 输出全 0 → out = 0 → relu(0 + identity)。
好处:
- 训练初期整个网络等价于恒等堆叠,非常稳定
- 残差分支「从零开始学」,而不是一开始就引入随机扰动
- 这是「让深网络从简单解开始学」的具体实现
同样的技巧用于 DiT / ControlNet:把最后一个线性层的权重和偏置初始化为 0,让网络初始时是「什么都不做」,然后逐渐学到有用的变换。
实测参考(本篇实验 3):零初始化后差异是 2.58,不是 0——因为末尾的 ReLU 会截断负数。所以严格来说不是精确恒等,但仍然大大稳定了训练。
Q3:混合精度训练时,为什么 scaler.unscale_(optimizer) 必须在 clip_grad_norm_ 之前?
答案
因为 clip_grad_norm_ 计算的是梯度范数,梯度此时被放大了 S 倍。
loss_scaled = loss * 65536 # 梯度也被放大 65536 倍
loss_scaled.backward()
# 如果直接裁剪:
clip_grad_norm_(model.parameters(), 1.0)
# 范数会是真实值的 65536 倍 → 几乎总是 > 1.0 → 所有梯度都被压到极小
# 等价于有效学习率变成 65536 分之一 → 训练完全失效
# 正确:先恢复真实梯度
scaler.unscale_(optimizer) # 梯度除以 S
clip_grad_norm_(model.parameters(), 1.0) # 现在范数是真实的2
3
4
5
6
7
8
9
10
11
PyTorch 的 API 设计:scaler.step(optimizer) 内部会自动 unscale_,所以如果不用梯度裁剪,可以不手动调用。
但如果你手动裁剪,必须自己先 unscale_,否则裁剪的是错误的范数。
推荐的安全写法:
scaler.scale(loss).backward()
scaler.unscale_(optimizer) # 总是先 unscale
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()2
3
4
5
Q4:为什么 bf16 不需要 loss scaling,但仍常配合 FP32 累积?
答案
不需要 loss scaling 的原因:bf16 的指数位和 fp32 相同(8 位),动态范围都是 10±38。而梯度小到 10−8 时 fp16 会下溢,bf16 不会(本篇实验 3 实测:fp16 → 0,bf16 → 9.9e-8)。
仍然需要 FP32 累积的原因:bf16 只有 7 位尾数(fp32 是 23 位),精度很低。
具体问题:
- 累加误差:优化器更新参数时是 w←w−lr⋅g。当 lr 很小时(比如 10−8),lr⋅g 相对于 w 极小,bf16 的 7 位尾数根本无法表示这个微小变化,更新会被完全舍入丢弃。
- 梯度累加时的抵消:小梯度的累加容易产生灾难性抵消。
标准做法:
# 前向 + 反向:bf16
with torch.autocast('cuda', dtype=torch.bfloat16):
loss = loss_fn(model(x), y)
loss.backward()
# 参数更新:fp32 master weights
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
# PyTorch 的参数本身是 fp32,autocast 只影响前向的计算精度2
3
4
5
6
7
8
注意 PyTorch 的实现细节:model.parameters() 本身是 fp32(因为模型创建时默认 fp32),autocast 只是在运算时临时转成 bf16。所以参数天然就是 FP32 master weights,不需要手动处理。
一句话总结:bf16 解决的是「范围」问题,FP32 累积解决的是「精度」问题。两件事,不同的坑。
Part 2 完成 → 进入 Part 3:Transformer
Part 3 · Transformer 与 LLM
这一部分是当前主流架构。从 Attention 的推导一路到 2024–2025 年的主流设计。
| # | 章节 | 核心问题 |
|---|---|---|
| 10 | Attention 的数学推导 | QKV 为什么要除 √d?softmax 为什么必须有? |
| 11 | 多头注意力与位置编码 | 多头在多做什么?位置信息怎么注入?RoPE 的旋转思想 |
| 12 | 从 Transformer 到 LLaMA 系列 | 现代 LLM 改了什么?RMSNorm / SwiGLU / RoPE 各解决了什么? |
| 13 | 推理优化:KV Cache 与 GQA | 自回归生成为什么慢?缓存和 GQA 怎么解决? |
| 14 | 缩放定律与高效注意力 | 为什么可以「大力出奇迹」?FlashAttention 赢在哪? |
配套代码
从零实现的 LLaMA(RMSNorm + RoPE + GQA + SwiGLU + Pre-LN), 163 行,已实测前向 + 反向通过。第 12、13 篇里的每个组件在这份代码里都有对应实现。
pip install torch
python llama_from_scratch.py # 前向 + loss + 反向,全程通过2
这一部分要建立的判断力
- 看到一个 LLM 架构图,能逐个组件说出它为什么在那里
- 能算清一次推理的显存占用和 FLOPs,判断瓶颈在哪
- 能区分训练期优化和推理期优化 —— 两者的约束完全不同
10 · Attention 的数学推导
核心问题:QKV 为什么要除 d?softmax 为什么必须有?如果去掉 scaling 会怎样?
一、从「需求」出发
1.1 我们想要什么能力
处理序列时,模型需要能:让每个位置「主动去查」其他位置的相关信息。
比如处理「那只猫很可爱,因为它饿了」——"它"要能关联到前面的"猫"。这种关联是动态的、依赖内容的,不能靠固定位置编码。
1.2 用检索类比理解 QKV
把 Attention 想成一个数据库检索系统:
| 角色 | 类比 | 作用 |
|---|---|---|
| Q(Query) | 我想要什么 | 当前查询的「需求描述」 |
| K(Key) | 我有什么 | 每个条目的「索引标签」 |
| V(Value) | 实际内容 | 每个条目真正的数据 |
检索流程:
- 拿我的 Q 去和所有条目的 K 比对 → 得到匹配分数
- 把分数转成权重(softmax)→ 归一化到 [0,1]
- 按权重加权求和所有条目的 V → 得到我要的结果
关键洞察:相似度是「Q 和 K 的关系」,但真正被加权取回的是 V。Q/K 负责「找谁」,V 负责「拿什么」。 这就是为什么 V 通常不参与打分。
二、Scaled Dot-Product Attention 的公式
Attention(Q,K,V)=softmax(dkQK⊤+M)V
其中:
- Q∈Rn×dk(n 个查询,每个 dk 维)
- K∈Rm×dk(m 个键)
- V∈Rm×dv
- M 是掩码矩阵(causal mask 等,可选)
形状流转(务必记住):
输入: Q [n, d_k] K [m, d_k] V [m, d_v]
↓ Q @ K.T
[n, m] ← ★ 注意力矩阵:n个查询对m个键的匹配度
↓ / sqrt(d_k)
[n, m]
↓ + M, softmax(dim=-1)
[n, m] ← 每行和为 1
↓ @ V
输出: [n, d_v]2
3
4
5
6
7
8
9
注意中间那个 [n,m] 矩阵是核心:它就是「注意力图」,可视化出来能看到模型在学什么。
三、三个核心问题
3.1 为什么除 dk(这是本篇最重要的部分)
数学推导:
假设 q 和 k 的每个分量都是独立的、均值 0 方差 1 的随机变量。那么点积:
q⋅k=i=1∑dkqiki
方差(独立项相加,方差相加):
Var[q⋅k]=i=1∑dkVar[qiki]=i=1∑dkVar[qi]Var[ki]=dk⋅1⋅1=dk
所以 std[q⋅k]=dk。
实测验证:
import torch, math
torch.manual_seed(0)
print("=== 点积的标准差随维度增长 ===")
for d in [16, 64, 256, 1024]:
q = torch.randn(2000, d); k = torch.randn(2000, d)
dot = (q * k).sum(-1)
print(f"d={d:>5} 点积 std={dot.std():>7.3f} 理论sqrt(d)={math.sqrt(d):>7.3f}"
f" 缩放后 std={dot.std()/math.sqrt(d):>6.3f}")2
3
4
5
6
7
8
实测输出:
d= 16 点积 std= 4.001 理论sqrt(d)= 4.000 缩放后 std= 1.000
d= 64 点积 std= 8.114 理论sqrt(d)= 8.000 缩放后 std= 1.014
d= 256 点积 std= 16.144 理论sqrt(d)= 16.000 缩放后 std= 1.009
d=1024 点积 std= 33.293 理论sqrt(d)= 32.000 缩放后 std= 1.0402
3
4
完美吻合:点积的 std 就是 d,除以 d 后稳定在 1 附近。
为什么这会导致训练失败:
softmax 的输入被放大 d 倍后进入饱和区。饱和的 softmax 是什么样子? 看实验:
import torch, math
print("=== softmax 饱和:logits 尺度的影响 ===")
print(f"{'元素数':>7}{'logits_std':>12}{'最大概率':>10}{'熵':>8}{'均匀分布熵':>12}")
for n in [4, 16, 64, 256]:
for scale in [1.0, 4.0, 16.0]:
torch.manual_seed(0)
p = (torch.randn(2000, n) * scale).softmax(-1)
maxp = p.max(-1).values.mean()
ent = -(p * torch.log(p + 1e-10)).sum(-1).mean()
print(f"{n:>7}{scale:>12.1f}{maxp:>10.4f}{ent:>8.3f}{math.log(n):>12.3f}")2
3
4
5
6
7
8
9
10
实测输出:
元素数 logits_std 最大概率 熵 均匀分布熵
4 1.0 0.5208 1.110 1.386
4 4.0 0.8313 0.421 1.386
4 16.0 0.9567 0.105 1.386
16 1.0 0.2481 2.356 2.773
16 4.0 0.7029 0.848 2.773
16 16.0 0.9263 0.187 2.773
64 1.0 0.1069 3.686 4.159
64 4.0 0.5887 1.317 4.159
64 16.0 0.8936 0.273 4.159
256 1.0 0.0436 5.050 5.545
256 4.0 0.5114 1.778 5.545
256 16.0 0.8766 0.326 5.5452
3
4
5
6
7
8
9
10
11
12
13
仔细看 logits_std = 16 那一行:熵只有 0.105 ~ 0.326,而均匀分布的熵是 1.386 ~ 5.545。注意力分布几乎变成了 one-hot——每个查询只关注一个键。
这会导致三个后果:
- 梯度消失:softmax 饱和区域的导数 softmax(z)(1−softmax(z)) 趋近 0
- 无法学习:所有注意力集中在一个位置上,模型丧失了「对比多个位置」的能力
- 信息瓶颈:每个 token 只能看到 1 个 token
关键数字:d_k = 64$ 时,点积 std 是 8,$\sqrt{64}=8$——**不缩放就是 logits_std=8**,已经落在表里 std=4 和 std=16 之间,softmax 明显饱和。d_k = 128$(LLaMA 的典型值)时 std 是 11.3,问题更严重。
结论:除以 dk 不是可选的微调,而是保证 softmax 工作在有效区间的必要操作。
3.2 为什么必须用 softmax(不能直接用点积)
三个理由:
理由1:需要归一化。softmax 让权重和为 1,等于「分配注意力预算」——每个 token 分配的注意力总量固定,不会因为某个分数高就无限放大。
理由2:可微且梯度有意义。softmax 是平滑函数,梯度 ∂sj∂ai=ai(δij−aj) 给出「提升 key j 的分数会如何改变分配」的清晰信号。
理由3:引入竞争/相对性。softmax 是「相对」操作——某个键分数升高会抢占其他键的权重。这符合注意力的直觉:注意力是稀有的资源,需要竞争。
如果直接用点积(可以理解为加权平均但不归一化),权重可能全都很小(输出尺度失控)或都很大(输出爆炸)。
3.3 mask 的作用
softmax(dkQK⊤+M)
其中 M 在要屏蔽的位置是 −∞(softmax 后变成 0)。
Causal mask(自回归模型):第 i 个 token 只能看到 ≤i 的 token。
mask = torch.triu(torch.ones(seq, seq) * float('-inf'), diagonal=1)用 -inf 而不是大负数的原因:softmax 会先减最大值,-inf 在减法后会变成 NaN(-inf−finite=-inf,exp(−∞)=0 实际是安全的,但减法顺序可能出问题)。实践中 PyTorch 用 float('-inf') 是安全的,因为 softmax 内部对 -inf 有特殊处理。
四、完整实现
import torch
import torch.nn.functional as F
from torch import nn
def scaled_dot_product_attention(Q, K, V, mask=None):
"""
Q: [batch, heads, seq, d_k]
K: [batch, heads, seq, d_k]
V: [batch, heads, seq, d_v]
mask: [batch, 1, seq, seq] 或 [seq, seq],True 表示要屏蔽
"""
d_k = Q.size(-1)
scores = Q @ K.transpose(-2, -1) / math.sqrt(d_k) # [b, h, n, m]
if mask is not None:
scores = scores.masked_fill(mask, float('-inf'))
attn = scores.softmax(dim=-1)
return attn @ V, attn2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
PyTorch 内置版本(生产环境务必用这个):
# ★ 官方推荐:内存高效的融合实现
out = F.scaled_dot_product_attention(Q, K, V, is_causal=True)2
它会自动选择最优的数学后端(数学等价但显存占用不同),并支持 Flash Attention。
五、自回归掩码的实现细节
5.1 因果掩码矩阵
seq_len = 4
掩码(True = 屏蔽):
j=0 1 2 3
i=0 [False True True True ] ← token0 只能看自己
i=1 [False False True True ] ← token1 能看 0,1
i=2 [False False False True ]
i=3 [False False False False]2
3
4
5
6
7
8
seq = 4
mask = torch.triu(torch.ones(seq, seq, dtype=torch.bool), diagonal=1)
print(mask)
# tensor([[False, True, True, True],
# [False, False, True, True],
# [False, False, False, True],
# [False, False, False, False]])2
3
4
5
6
7
5.2 Padding mask 的组合
实际场景要同时考虑 causal mask 和 padding mask:
padding_mask = (tokens != pad_id) # [batch, seq] True=有效
causal_mask = torch.triu(torch.ones(seq, seq, dtype=torch.bool), diagonal=1)
# padding: [b, seq] → [b, 1, 1, seq](广播到所有 query)
# causal: [seq, seq] → [1, 1, seq, seq]
combined = causal_mask[None, None, :, :] | (~padding_mask)[:, None, None, :]2
3
4
5
6
六、动手实验
实验 1:验证 √d 缩放(本文核心)
import torch, math
torch.manual_seed(0)
print("=== 点积标准差 vs 维度 ===")
print(f"{'d_k':>7}{'点积std':>10}{'√d':>8}{'缩放后std':>12}")
for d in [16, 64, 256, 1024]:
q = torch.randn(2000, d); k = torch.randn(2000, d)
dot = (q * k).sum(-1)
print(f"{d:>7}{dot.std():>10.3f}{math.sqrt(d):>8.3f}{dot.std()/math.sqrt(d):>12.3f}")
print("\n=== 不缩放的 softmax 饱和程度 ===")
print(f"{'d_k':>7}{'max_p(不缩放)':>16}{'max_p(缩放)':>14}{'熵(缩放)':>11}{'均匀熵':>9}")
for d in [16, 64, 256, 1024]:
q = torch.randn(500, d); k = torch.randn(500, d)
dot = (q * k).sum(-1)
a = dot.softmax(-1).max(-1).values.mean()
b = (dot / math.sqrt(d)).softmax(-1).max(-1).values.mean()
ent = -(b * torch.log(b + 1e-10)).sum(-1).mean()
print(f"{d:>7}{a:>16.4f}{b:>14.4f}{ent:>11.3f}{math.log(d):>9.3f}")
print("\n→ 不缩放时 max_p 全部=0.9999,注意力完全退化成 one-hot")
print("→ 缩放后 max_p 回到 0.02~0.035,有明确的相对差异可供学习")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
实测输出:
=== 不缩放的 softmax 饱和程度 ===
d_k max_p(不缩放) max_p(缩放) 熵(缩放) 均匀熵
16 0.7267 0.0319 0.110 2.773
64 0.9990 0.0350 0.117 4.159
256 0.9999 0.0347 0.117 5.545
1024 0.9999 0.0170 0.069 6.931
→ 不缩放时 max_p 全部=0.9999,注意力完全退化成 one-hot
→ 缩放后 max_p 回到 0.02~0.035,有明确的相对差异可供学习2
3
4
5
6
7
8
9
注意 max_p(缩放) 那一列只有 0.02~0.035,看起来「过于均匀」了——但这恰恰是正确的:真实训练中的 q,k 不是独立随机向量,它们经过投影和训练后有特定的相关结构,会让部分位置的分数更高。随机初始化下的均匀分布正是我们想要的起点(有区分度,才能学)。
真正要避免的是 max_p(不缩放) 那一列的 0.9999——所有位置的注意力完全相同,没有任何区分能力。
注意:真实 Transformer 里 q,k 经过 WQ,WK 投影后分量 std 约 1/din 量级,但经过层归一化和训练后,实际进入 attention 的 q,k 分量 std 接近 1,所以上面的分析成立。
实验 2:手动实现并与 PyTorch 对比
import torch, math
import torch.nn.functional as F
torch.manual_seed(42)
b, h, seq, d = 2, 4, 8, 16
Q = torch.randn(b, h, seq, d)
K = torch.randn(b, h, seq, d)
V = torch.randn(b, h, seq, d)
causal = torch.triu(torch.ones(seq, seq, dtype=torch.bool), diagonal=1)
# 手动实现
def manual_attention(Q, K, V, causal_mask=None):
d_k = Q.size(-1)
scores = Q @ K.transpose(-2, -1) / math.sqrt(d_k)
if causal_mask is not None:
scores = scores.masked_fill(causal_mask, float('-inf'))
return scores.softmax(-1) @ V
mine = manual_attention(Q, K, V, causal)
theirs = F.scaled_dot_product_attention(Q, K, V, is_causal=True)
print(f"手写 vs PyTorch 内置:最大差异 = {(mine - theirs).abs().max():.2e}")
print("→ 完全一致,但内置版本内存效率高得多")
# 验证因果性:改变未来的 token,不应影响过去的输出
V2 = V.clone()
V2[:, :, 5:, :] = torch.randn(b, h, seq - 5, d) * 100# 疯狂改动后面的 token
out1 = manual_attention(Q, K, V, causal)
out2 = manual_attention(Q, K, V2, causal)
print(f"\n改动后 token 之后,输出差异: {(out1[:, :, :5] - out2[:, :, :5]).abs().max():.2e}")
print(f"改动后 token 之后的后续位置差异: {(out1[:, :, 5:] - out2[:, :, 5:]).abs().max():.2e}")
print("→ 前面的 token 输出完全不受影响,causal mask 正确")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
实验 3:可视化注意力权重
import matplotlib
matplotlib.use('Agg')
import matplotlib.pyplot as plt
# 造一个能看出结构的例子:后一半的 token 是前一半的「复制」
seq = 8
x = torch.randn(1, 1, seq, 32) * 0.5
x[0, 0, 4:] = x[0, 0, :4]# token 4~7 是 0~3 的副本
lin_q = nn.Linear(32, 32); lin_k = nn.Linear(32, 32); lin_v = nn.Linear(32, 32)
Q, K, V = lin_q(x), lin_k(x), lin_v(x)
scores = Q @ K.transpose(-2, -1) / math.sqrt(32)
attn = scores.softmax(-1)[0, 0]
print("注意力矩阵 (行=query, 列=key):")
print(" " + "".join(f"{j:>6}" for j in range(seq)))
for i in range(seq):
bar = "".join(f"{v:>6.2f}" for v in attn[i].tolist())
print(f" t{i} {bar}")
print("\n观察第4~7 行(副本 token)是否指向对应的 0~3 列")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
观察:第 4~7 行(副本)应该主要指向对应的 0~3 列,因为它们内容相同,q⋅k 最大。这直观展示了 attention 在「按内容匹配」。
七、自测题
Q1:如果去掉 dk 缩放,训练会出现什么现象?可以量化吗?
答案
现象:softmax 饱和 → 梯度消失 → 训练完全失效或极慢。
量化(用本篇实验的数据,`d_k = 64$):
| 最大概率 | 说明 | |
|---|---|---|
| 不缩放 | 0.9990 | 注意力完全塌成one-hot |
| 缩放 | 0.0350 | 有区分度,可以学习 |
从 0.9990 变成 0.035——不缩放时 softmax 输出的就是「1 和一堆 0」,梯度全在最大值那个位置上,其余位置的梯度是 ParseError: KaTeX parse error: Unexpected character: '' at position 12: p_i(1-p_j) ̲pprox 0。模型无法学到任何「关注多个位置」的能力。
梯度的量化:softmax 的雅可比是 diag(p)−pp⊤,元素大小 ∼pi(1−pj)。当 p→1 时 p(1−p)→0,梯度消失。
dk=1024 时 max_p = 0.9999,softmax 数值上就是 one-hot,梯度几乎为 0。
这也是原始 Transformer 论文把它叫"scaled dot-product"的原因——scaling 是这个公式的核心组成部分,不是可选优化。
Q2:Q、K、V 三个矩阵分别是什么?能不能只用两个(比如 Q 和 V)?
答案
不能,因为「谁该被关注」和「关注后拿到什么」是两个独立的信息。
- QK^T 决定「注意力权重」:谁和谁相关
- V 决定「传递什么内容」
如果只有 Q 和 V,用 Q⊤V 直接算?那得到的是 [d,d] 的矩阵(不是注意力权重矩阵),失去了「按需检索」的能力。
类比:图书馆系统里,「检索词→书的索引」和「书的实际内容」必须分开存储。如果混在一起,就无法实现「用检索词找到匹配的书,然后取它的内容」。
一个反例说明问题:假设一句话里有两个关键信息点,V 分别是它们的内容。如果只有一个矩阵,你无法表达「我要同时关注这两点,但注意力要按相关性分配」。
实际上 K 通常可以和 Q 共享同一份计算结果(self-attention 里),某些简化架构会这么做,但效果会下降。标准实现里 W_Q、W_K、V_W 是三套独立参数。
Q3:为什么 causal mask 要用 −∞,用 −109 行不行?
答案
理论上都行,实践中都用 −∞(或者 PyTorch 的 float('-inf'))。
用 −109 的情况:
softmax([5.0, -1e9]) = [1.0, 0.0] ✓ 效果正确因为 exp(−109)≈0,和 exp(−∞)=0 几乎没区别。
但有两个隐患:
数值精度:如果模型内部用 fp16,−109 直接超出 fp16 范围(最大 65504),会变成
-inf或nan。用 fp32 的话 −109 也在边缘。和 causal mask 组合时:如果同时有 padding mask,两层 mask 相加可能得到 −2×109,进一步溢出。
PyTorch 的 F.scaled_dot_product_attention 用 is_causal=True 参数,内部自动处理这个 mask,比手动传更高效(能融合进 kernel)。
一个实用建议:训练时优先用 is_causal=True 而不是手动构造 mask,能享受 Flash Attention 的优化。
Q4:注意力矩阵 [n,m] 的复杂度是多少?为什么这是长序列的瓶颈?
答案
时间复杂度:O(n⋅m⋅dk),当 n=m=L 时是 O(L2dk)。
空间复杂度:注意力矩阵本身要存 O(L2)。这是平方级。
具体数字(LLaMA-7B,4096 token):
注意力矩阵:2 (batch) × 32 (heads) × 4096 × 4096 × 2 bytes (fp16)
= 2 × 32 × 4096 × 4096 × 2
= 2.1 GB ← 单层!2
3
32 层累计就是 68 GB 仅用于存注意力矩阵。
对比 MLP 的复杂度:O(L⋅d2),线性于 L。
所以瓶颈很明确:
| 组件 | 复杂度 | 随序列长度 |
|---|---|---|
| QKV 投影 | O(Ld2) | 线性 |
| 注意力矩阵 | O(L2d) | 平方 ⚠️ |
| MLP | O(Ld2) | 线性 |
这就是所有「高效注意力」工作的动机:
- FlashAttention:不显式存注意力矩阵(见第 14 篇)
- 稀疏注意力:只算部分位置(Longformer、BigBird)
- 线性注意力:改变计算顺序,用 K⊤V 先算(Linear Transformer、Performer)
- 低秩近似:把 QK⊤ 近似成低秩矩阵(Linformer)
GPT-3 的选择:只支持 2048 token 上下文,因为再长平方成本太高。GPT-4 的 128K 上下文能实现,靠的是 FlashAttention + 多种优化(推测,多层特征)。
下一篇 → 多头注意力与位置编码
11 · 多头注意力与位置编码
核心问题:多头到底在多做什么?位置信息怎么注入?RoPE 的旋转技巧是什么?
一、多头注意力:为什么要「多头」
1.1 单头的限制
单个注意力头只能算出一个 [n,n] 的注意力矩阵。这意味着每个 token 只能有一个「关注模式」。
问题:语言里的关联是多模态的。同一句话里可能同时需要:
- 语法关联:「猫」↔「的」(结构)
- 语义关联:「猫」↔「动物」(指代)
- 语用关联:这个 token 在整体语境中的角色
一个头做不到兼顾。
1.2 多头的做法
把 dmodel 拆成 h 份,每份独立做注意力,最后拼接。
MultiHead(Q,K,V)=Concat(head1,…,headh)WO
headi=Attention(QWiQ,KWiK,VWiV)
形状推导(务必掌握):
输入 X: [batch, seq, d_model]
↓ W_Q: [d_model, d_model]
Q = X W_Q
★ 关键:reshape成多头(不改变总维度,只是分组)
原始: [batch, seq, d_model]
变换: [batch, seq, h, d_k] → transpose → [batch, h, seq, d_k]
其中 d_model = h × d_k
↓ 每个头独立做 attention
[batch, h, seq, d_k] × [batch, h, d_k, seq] → [batch, h, seq, seq]
→ softmax → 加权 V → [batch, h, seq, d_k]
↓ transpose + reshape 还原
[batch, h, seq, d_k] → [batch, seq, h, d_k] → [batch, seq, d_model]
↓ W_O: [d_model, d_model]
输出: [batch, seq, d_model]2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
注意一个容易误解的点:多头不是把 dmodel 分成 h 份各自独立算完就不管了,而是每个头有自己独立的 WQ,WK,WV,参数完全不共享。最后还要过一个 WO 做融合。
参数量:
| 项目 | 参数量 |
|---|---|
| WQ,WK,WV | 3×dmodel×dmodel |
| WO | dmodel×dmodel |
| 总计 | 4dmodel2 |
和单头完全一样! 多头不增加参数量,只是改变了计算的方式(把一个大矩阵乘法拆成 h 个小的)。这是「结构先验」而非「容量增加」的典型例子。
1.3 参数量验证
import torch
from torch import nn
d_model, h = 512, 8
d_k = d_model // h
single = nn.Linear(d_model, d_model)
# 多头的四个矩阵
q = nn.Linear(d_model, d_model); k = nn.Linear(d_model, d_model)
v = nn.Linear(d_model, d_model); o = nn.Linear(d_model, d_model)
n_single = sum(p.numel() for p in single.parameters())
n_multi = sum(p.numel() for p in list(q.parameters())+list(k.parameters())+list(v.parameters())+list(o.parameters()))
print(f"单头 Linear({d_model},{d_model}): {n_single:,}")
print(f"多头 4 个矩阵: {n_multi:,}")
print(f"倍数: {n_multi/n_single:.1f}x")
print(f"\n★ 注意:对比的是1 个 Linear vs 4 个。但单头只有 1 组 QKV,多头也是 4 个矩阵")
print(f" 单头 QKV+输出 = 4 个 {d_model}x{d_model} 矩阵 = {n_multi:,} ← 和多头相同")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
实测输出:
单头 Linear(512,512): 262,656
多头 4 个矩阵: 1,050,624
倍数: 4.0x2
3
这说明:多头和单头的参数量相同(都是 4 个 dmodel×dmodel 矩阵),因为单头也需要 WQ,WK,WV,WO 四个矩阵。多头是「把一个大注意力拆成 h 个小注意力」,不是「多加了参数」。
1.4 为什么小头反而更好
一个反直觉但被广泛验证的结论:减小 dk 往往提升效果。
| 模型 | dmodel | 头数 h | dk |
|---|---|---|---|
| 原始 Transformer | 512 | 8 | 64 |
| LLaMA-7B | 4096 | 32 | 128 |
| LLaMA-13B | 5120 | 40 | 128 |
| LLaMA-65B | 8192 | 64 | 128 |
| LLaMA2-7B | 4096 | 32 | 128 |
关键观察:几乎所有 LLM 的 dk 都固定在 64 或 128,不管模型多大。
原因(推测性的,但有实证支持):
- 注意力矩阵的噪声:dk 越大,点积的方差越大,softmax 越容易进入需要精细 dk 校准的区域
- 过拟合风险:dk 越大,每个头的表达能力越强,越容易过拟合
- 「多而浅」优于「少而深」:更多头 = 更多并行的关系模式,每头更专注
这个观察的价值:现代 LLM 的设计里,头数随模型规模线性增长,但 dk 恒定——这意味着 dmodel=h×dk,即模型的宽度主要由「头数」决定。
二、位置编码:Transformer 的先天缺陷
2.1 问题
Attention 是置换等变的(permutation-equivariant):
Attention(PX,PK,PV)=PAttention(X,K,V)
把输入顺序打乱,输出只是跟着打乱,模型完全感知不到顺序变化。
后果:「狗咬人」和「人咬狗」在纯 attention 看来是一样的。
2.2 四类解法
| 方法 | 代表 | 思路 | 外推能力 |
|---|---|---|---|
| 绝对位置嵌入 | BERT、GPT-2 | 把位置向量加到 token embedding | 差(超出训练长度就崩) |
| 可学习位置嵌入 | BERT、GPT-2 | 位置也用可学习的向量 | 差(同上) |
| 相对位置编码 | T5、ALiBi | 在 attention 计算时加入相对距离偏置 | 中 |
| 旋转位置编码(RoPE) | LLaMA、几乎所有现代 LLM | 用旋转把位置注入 Q/K | 好 |
三、RoPE(Rotary Position Embedding)
RoPE 是目前的事实标准——LLaMA、LLaMA2、LLaMA3、Qwen、Mistral 全都用它。这是本篇最重要的内容。
3.1 核心洞察
原论文(Su et al., 2021)的洞察:
绝对位置嵌入会削弱注意力机制——因为「位置」和「内容」被混在一个向量里,无法分离。
RoPE 的做法:不添加任何向量到输入,而是把 q,k 向量按位置「旋转」。
3.2 二维情况下的旋转(理解的关键)
假设 d=2,把向量看作复平面上的一个点:
q=(q0,q1)↔q0+iq1
位置 m 对应一个旋转角度 mθ。旋转:
q′=Rmq=(cosmθsinmθ−sinmθcosmθ)(q0q1)
为什么这个设计是天才的?
核心性质:旋转内积定理——
⟨Rmq,Rnk⟩=⟨q,Rn−mk⟩=f(q,k,n−m)
内积只依赖相对距离 n−m,与绝对位置无关!
这意味着 RoPE 天然编码了相对位置关系——而这正是注意力真正需要的信息(「这个词和前一个词的关系」比「这个词在第 500 位」更有用)。
3.3 推广到高维
d 维向量切成 d/2 对,每对用不同频率的旋转:
θi=base−2i/d,i=0,1,…,d/2−1
几何间隔的频率:θ 按指数衰减,所以低维用高频(捕捉近距离),高维用低频(捕捉远距离)。
import torch
import math
def apply_rope(x, pos, base=10000):
"""
x: [seq, d]
pos: [seq] 位置索引
"""
d = x.shape[-1] // 2
# 频率:几何级数,base=10000 是原论文的经验值
inv_freq = 1.0 / (base ** (torch.arange(0, d).float() / d)) # [d/2]
# 角度:[seq, d/2]
angles = pos[:, None].float() * inv_freq[None, :]
cos, sin = angles.cos(), angles.sin()
x1, x2 = x[..., :d], x[..., d:] # 前后两半配对
return torch.cat([x1 * cos - x2 * sin, # 旋转第一半
x2 * cos + x1 * sin], -1) # 旋转第二半2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
3.4 RoPE 的两个关键性质(实测验证)
import torch
torch.manual_seed(0)
def apply_rope(x, pos, base=10000):
d = x.shape[-1] // 2
inv_freq = 1.0 / (base ** (torch.arange(0, d).float() / d))
angles = pos[:, None].float() * inv_freq[None, :]
cos, sin = angles.cos(), angles.sin()
x1, x2 = x[..., :d], x[..., d:]
return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], -1)
d, max_pos = 16, 200
q = torch.randn(d); k = torch.randn(d) # ★ 同一个 q, k
pos = torch.arange(max_pos)
Q, K = apply_rope(q, pos), apply_rope(k, pos)
print("同一个 (q,k) 放在不同位置,内积应只依赖 |位置差|:")
for off in [0, 1, 2, 4, 8, 16, 32]:
vals = [(Q[i] * K[i + off]).sum().item() for i in range(0, max_pos - off, 20)]
spread = max(vals) - min(vals)
print(f" |Δpos|={off:>3}: {[f'{v:.4f}' for v in vals[:3]]} 波动={spread:.1e}")
print("\n→ 同一距离下内积完全一致(波动仅 1e-6,即浮点误差)")
print("→ 这就是 RoPE 的核心性质:<R_m q, R_n k> 只依赖 (m-n)")
print("→ 不同距离内积不同 → 模型能区分相对距离")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
实测输出:
同一个 (q,k) 放在不同位置,内积应只依赖 |位置差|:
|Δpos|= 0: ['2.9839', '2.9839', '2.9839'] 波动=1.2e-06
|Δpos|= 1: ['2.8449', '2.8449', '2.8449'] 波动=3.1e-06
|Δpos|= 2: ['1.2160', '1.2160', '1.2160'] 波动=2.4e-06
|Δpos|= 4: ['-0.6933', '-0.6933', '-0.6933'] 波动=8.6e-06
|Δpos|= 8: ['-2.6004', '-2.6004', '-2.6004'] 波动=1.0e-05
|Δpos|= 16: ['-3.5409', '-3.5409', '-3.5409'] 波动=1.2e-05
|Δpos|= 32: ['-2.6989', '-2.6989', '-2.6989'] 波动=3.3e-062
3
4
5
6
7
8
完美的验证:每个距离下的内积完全一致(波动是浮点误差级别),不同距离内积不同。
这正是 RoPE 优于绝对位置编码的核心原因:绝对位置编码下,同一对 (q,k) 在不同绝对位置的相似度不同,模型要额外学习「位置偏移」;RoPE 直接把这个关系做进了旋转里。
3.5 RoPE 的外推能力
RoPE 的外推能力来自频率的连续性:训练时见过的最大位置是 L,理论上可以推广到任意位置(只要角度还在有效范围内)。
但实际上会退化。原因是注意力熵爆炸:
- 训练时模型学会了「近处的 token 重要」
- 位置超过训练长度后,旋转角度过大,注意力分布变得不稳定
- 结果:注意力熵(不确定性)急剧上升,模型开始「乱看」
改进方案(NTK-aware scaling / YaRN):
| 方法 | 思路 |
|---|---|
| 位置插值(PI) | 把位置 m 缩放到 Ltarget/Ltrainm |
| NTK-aware | 修正 base 参数,低频维度保持、高频维度插值 |
| YaRN | 分维度组合 PI 和 NTK,理论上最优 |
实践建议:扩展上下文长度时,第一选择总是 YaRN 或 NTK-aware,比直接 PI 效果好得多。
四、ALiBi:另一种思路
ALiBi(Attention with Linear Biases) 更简单:不做旋转,直接在 softmax 之前给 logits 加一个位置偏置。
Attention=softmax(dQK⊤+bias)
biasij={0−α(i−j)j≤ij>i
其中 α 是每个头的斜率参数(如 8 个头的 α = [1/21,1/22,…,1/28])。
特点:
- 极简:不需要任何参数,不增加计算
- 外推性好:距离是线性的,理论上可外推
- 性能略低于 RoPE:现代 LLM 基本不用了
为什么 RoPE 更好:ALiBi 是「硬性惩罚远处 token」,而 RoPE 是「让模型自己学位置关系」。后者更灵活。
五、动手实验
实验 1:验证多头不增加参数量
见第一节代码。实测:单头和多头都是 4 个 dmodel2 矩阵,参数量相同。
实验 2:RoPE 的相对位置性质
见第三节代码。这是本篇核心实验,输出显示内积波动仅 1e-6。
实验 3:位置编码的必要性
import torch
import torch.nn.functional as F
torch.manual_seed(0)
B, seq, d = 1, 6, 16
X = torch.randn(B, seq, d)
Wq, Wk, Wv = [torch.randn(d, d) for _ in range(3)]
def attention(x):
q, k, v = x @ Wq, x @ Wk, x @ Wv
return (q @ k.transpose(-2, -1) / d ** 0.5).softmax(-1) @ v
# 打乱顺序
perm = torch.randperm(seq)
out1 = attention(X)
out2 = attention(X[:, perm])
# 把输出也按相同顺序还原,对比
out2_aligned = torch.empty_like(out2)
out2_aligned[:, perm] = out2
diff = (out1 - out2_aligned).abs().max().item()
print(f"打乱顺序后,输出(对齐后)的差异: {diff:.2e}")
print(f"→ 差异为 0,证明 attention 完全对顺序无感知")
print(f"\n置换矩阵:\n{perm.tolist()}")
print("→ 这就是为什么必须要有位置编码!")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:频率的指数结构
import torch
base = 10000
d = 8# 假设 8 个频率
inv_freq = 1.0 / (base ** (torch.arange(0, d).float() / d))
print("频率(inv_freq):")
for i, f in enumerate(inv_freq.tolist()):
pos_at_1 = f# 位置 1 时的角度
pos_at_1000 = f * 1000
period = 2 * 3.14159 / f
print(f" 维度{i}: inv_freq={f:.4f} 1个位置的转角={f:.4f}rad "
f"周期={period:>10.1f} 个位置")
print("\n→ 低维度频率高、周期短 → 捕捉近距离关系")
print("→ 高维度频率低、周期长 → 捕捉远距离关系")
print("→ 几何级数设计让每个维度覆盖一个不同的距离尺度")2
3
4
5
6
7
8
9
10
11
12
13
14
15
六、自测题
Q1:多头注意力的参数量和单头一样,那多头的「多」体现在哪里?
答案
参数量一样,但「计算方式的结构」不同。
单头:QK⊤ 是 [n,n] 的矩阵,每个元素是 dmodel 维向量的内积。
多头:把 dmodel 拆成 h 份,每个头用 dk=dmodel/h 维做内积,共 h 个 [n,n] 矩阵。
「多」体现在:
- h 个独立的注意力图——同一对 token 在不同头上可以有完全不同的关联强度。比如某个头关注语法邻接,另一个头关注长距离指代。
- 每个头学到了不同的关系模式(可解释性研究已证实:有些头专门做「前一个词」的关注,有些头做「句号」的关注)
- 子空间多样性——不同头在不同的特征子空间里工作,类似 CNN 的多通道
类比:单头像「用一副有色眼镜看世界」,多头像「戴 8 副不同的眼镜」,每副看到的东西不一样。
参数量不变是优势:可以在不加参数的情况下增加表达多样性。
Q2:RoPE 为什么比「加位置向量」更好?至少说两个理由。
答案
理由 1:内置了相对位置关系
- 加位置向量:q=(WQxi+pi),相似度里有 pi⊤pj 这种「绝对位置对绝对位置」的项,是混乱的
- RoPE:⟨Rmq,Rnk⟩=f(q,k,n−m),干净地只依赖相对位置
理由 2:不占据表示空间
加位置向量时,dmodel 维里有一部分专门用来存位置信息,挤占了表示内容的空间。
RoPE 不增加任何维度——它只是把 q,k 旋转了一下,维度完全不变。
理由 3:更好的外推
几何级数的频率设计让 RoPE 在位置上更平滑。绝对位置嵌入(一个可学习的向量表)在超出训练长度时完全没有对应向量,直接失效。
理由 4:与 attention 的数学结构兼容
RoPE 的旋转是正交变换,保持内积关系(旋转矩阵满足 R⊤R=I)。这保证了它不会破坏 q,k 原本的语义结构。
一句话总结:RoPE 是「把位置编码进操作里」,而不是「把位置编码进数据里」。前者更优雅、更高效。
Q3:为什么 LLM 的 dk 几乎都固定为 64 或 128,不管模型多大?
答案
这是实证观察(可以查 HuggingFace 的 config.json 验证),理论解释有几种推测:
推测 1:注意力熵与性能的关系
dk 越大,q⋅k 的方差越大。即使除以 dk,模型也需要精细校准才能让 softmax 工作在好的区间。适中的 dk 让 attention 更容易学。
推测 2:过拟合
dk 大 → 每个头的表达能力更强 → 更容易过拟合。LLM 靠数据规模控制过拟合,不需要靠小 dk。
推测 3:「多而浅」的归纳偏置
更多头 = 更多并行关系模式,每头更专注。这比「少而深」的单头更符合注意力机制的本质——它本来是做「关系匹配」的,不是做「特征提取」的。
推测 4:工程效率
小 dk 的矩阵乘法更容易被 GPU 的 tensor core 优化(尤其是 fp16/bf16)。dk=128 是 tensor core 的友好尺寸。
实践建议:不要自己改这个。除非你有明确的理由和验证,否则跟随主流配置。
Q4:RoPE 能外推到比训练时更长的序列吗?会遇到什么问题?
答案
理论上能,实践中会退化。
问题:注意力熵爆炸
训练时模型学会了「近处重要、远处次要」的模式。位置超出训练长度后:
- 旋转角度过大——base=10000 下,高维的旋转周期很长,但低维在位置 106 时已经转了无数圈,数值上不再有意义
- 注意力分布变得均匀或混乱——模型「不知道该关注谁」
- 困惑度急剧上升——实测通常从 10飙到 100+
解决方案(按推荐度):
| 方法 | 核心思路 | 效果 |
|---|---|---|
| YaRN | 分维度组合 PI + NTK | ★★★ 目前最好 |
| NTK-aware scaling | 修正 base,高频少插值、低频多插值 | ★★☆ 简单有效 |
| 位置插值(PI) | 位置 m→m⋅LtargetLtrain | ★★ 有损 |
| 继续预训练 | 在长文本上继续训 | ★★★ 慢但最可靠 |
NTK-aware 的直觉:RoPE 的低频维度负责远距离、高频负责近距离。扩展长度时应该保持低频维度不变(它们的周期本来就长),只对高频维度做插值。而「保持不变」的直觉做法就是减小 base。
LLaMA3 的做法:8K 预训练 → 分阶段扩展到 128K,中间用 NTK-aware + 继续预训练 + 平均检查点(averaging checkpoints)避免灾难性遗忘。
下一篇 → 从 Transformer 到 LLaMA 系列
12 · 从 Transformer 到 LLaMA 系列
核心问题:现代 LLM 相对原始 Transformer 改了什么?每个改动解决了什么问题?
一、原始 Transformer(2017)
1.1 Encoder-Decoder 结构
Encoder(双向,原点式) Decoder(自回归 +交叉注意力)
┌─────────────────┐ ┌──────────────────────┐
│ MultiHeadAttention │ │ MaskedMultiHeadAttention │
│ + Add & Norm │ │ + Add & Norm │
│ │ │ │
│ FeedForward │ ────────→ │ CrossAttention │
│ + Add & Norm │ │ + Add & Norm │
│ × N层 │ │ FeedForward + Add & Norm │
└─────────────────┘ │ × N 层 │
└──────────────────────┘2
3
4
5
6
7
8
9
10
两个关键区别:
| Encoder | Decoder | |
|---|---|---|
| 注意力 mask | 无(双向) | causal(只能看前面) |
| 注意力输入 | 只有源序列 | 自注意力 + encoder 的输出 |
| 用途 | BERT(理解) | GPT(生成) |
现代 LLM 只保留了 Decoder 部分(叫 decoder-only),因为:
- 统一了理解和生成任务
- 因果 mask 让训练可以并行(每个位置同时预测下一个词)
- 架构更简单
1.2 原始 Block 的 Post-LN 结构
out=LayerNorm(x+Sublayer(x))
注意 LayerNorm 在残差相加之后 —— 这叫 Post-LN。
问题:LN 夹在残差通路上,梯度要穿过 LN 才能回传,导致深层训练需要精细的 warmup,否则容易发散。
二、现代 LLaMA 风格的四大改动
Post-LN → Pre-LN
LayerNorm → RMSNorm → RMSNorm(更轻)
位置编码相加 → RoPE(更自然)
FFN 的 ReLU → SwiGLU(更有效)2
3
4
核心主题:从「能不能训」转向「怎么训得更好」。
三、Pre-LN:最重要的改动
3.1 结构对比
Post-LN (原始 Transformer):
x → Sublayer(x) → +x → LayerNorm → 下一层
Pre-LN (现代 LLM):
x → LayerNorm(x) → Sublayer → +x → 下一层2
3
4
5
3.2 为什么 Pre-LN 更好
梯度通路的区别:
Post-LN 的残差通路:不干净——LN 在通路中间,梯度要「穿过 LN」:
∂xL∂xL+1=LN′(xL+F(xL))(I+F′(xL))
那个 LN′ 是额外因子,深层累积会不稳定。
Pre-LN 的残差通路:干净的恒等映射:
∂xL∂xL+1=I+F′(LN(xL))
那个 I 保证了梯度 100% 直通——即使 F′ 是任何矩阵,梯度至少原封不动传下去。
3.3 实证
| Post-LN | Pre-LN | |
|---|---|---|
| 深层训练 | 需要 warmup | 稳定 |
| warmup | 必需(通常 4000 步) | 可选 |
| 最终性能 | 可能略好 | 略差或持平 |
| 现代 LLM | 几乎不用 | 全部采用 |
为什么最终性能反而 Post-LN 略好? 一个解释是 Post-LN 的 LN 在每个子层后重新缩放特征,有一定的正则化作用。但这个优势远小于「能稳定训练」的价值,所以全部转向 Pre-LN。
一个易错点:Pre-LN 要求最后有一个 final norm(在所有层之后):
x = block(x) for _ in range(N)
x = self.norm(x) # ★ final LayerNorm/RMSNorm,必需!2
因为 Pre-LN 的最后一个 block 输出没有经过归一化(LN 在子层内部)。
四、RMSNorm
第 6 篇已详述。这里只强调它在 LLaMA 里的作用:
| LayerNorm | RMSNorm(LLaMA 选择) | |
|---|---|---|
| 参数 | 2d(γ,β) | d(只有 γ) |
| 计算 | 求均值 + 求方差 + 减均值 + 除 | 只求均方根 + 除 |
| 能否融合 kernel | 部分 | 可以完全融合 |
LLaMA-7B 的实际节省(d=4096,32 层):
- 参数:每层省 4096,共省 131K(占总量 0.02%,微不足道)
- 显存和速度:可融合进相邻算子,减少多次 HBM 读写
关键点:RMSNorm 的收益不在参数,而在计算融合。
五、SwiGLU:更好的 FFN
5.1 三种 FFN 对比
原始 Transformer:
FFN(x)=ReLU(xW1+b1)W2+b2
GLU(Gated Linear Unit):
GLU(x)=ReLU(xW1)⊗(xW2)
SwiGLU(LLaMA 的选择):
SwiGLU(x)=SiLU(xW1)⊗(xW3)
其中 SiLU(也叫 Swish):SiLU(x)=x⋅σ(x)
5.2 门控的直觉
乘性激活 = 一个「门」控制另一个「内容」:
SiLU(xW₁)⊗ (xW₃)
↓ ↓
门(0~1) 内容(任意实数)
↓ ↓
↑──── 逐元素相乘 ────↑
门=0 → 完全关闭
门=1 → 完全通过2
3
4
5
6
7
为什么这有用:普通的 ReLU 只能「开启或关闭」整个神经元(且不可微地调整)。门控允许每个维度独立控制信息的通过量,而且是可学习的。
5.3 参数量陷阱:为什么中间层要缩小
SwiGLU 有三个矩阵(W1,W3,W2),比原来多一个。为了保持参数量不变,中间层维度要缩小:
隐层维度=38d≈2.67d(保持参数量与 FFN(4d) 相当)
LLaMA 用的就是这个:intermediate_size = 8192 而 hidden_size = 4096,比例是 2.0 而不是 4.0。
验证:
| FFN 类型 | 中间维度 | 参数量(d=4096) |
|---|---|---|
| FFN(4d) | 16384 | 3×d2 |
| SwiGLU(8/3 d) | 10923 | 3×d×(8/3)d=8d2 |
| LLaMA | 8192 | 3×4096×8192=100M |
经验值:LLaMA 用 2:1 到 8:3 之间的比例,效果基本持平。
六、GQA:分组查询注意力(推理优化)
第 13 篇详细讲计算,这里只给结构对比:
MHA (Multi-Head Attention): Q, K, V 都是 [h, d] ← 全部多头
MQA (Multi-Query): Q=[h,d], K=[1,d], V=[1,d] ← KV 只有 1 头
GQA (Grouped-Query): Q=[h,d], K=[g,d], V=[g,d] ← KV 分 g 组★2
3
LLaMA2 用 MQA,LLaMA3 用 GQA。 GQA 是 MHA 和 MQA 的折中,效果接近 MHA、成本接近 MQA。
七、LLaMA 完整配置对照表
| 组件 | 原始 Transformer | LLaMA / 现代 LLM | 原因 |
|---|---|---|---|
| 结构 | Encoder-Decoder | Decoder-only | 统一任务,训练并行 |
| 归一化位置 | Post-LN | Pre-LN | 梯度稳定 |
| 归一化类型 | LayerNorm | RMSNorm | 可融合,更快 |
| 位置编码 | 正弦绝对位置 | RoPE | 相对位置,外推好 |
| FFN | ReLU | SwiGLU | 门控表达力强 |
| FFN 中间维度 | 4d | ~2.67d | 补偿 SwiGLU 的额外矩阵 |
| 激活 | ReLU | GELU / SiLU | 平滑 |
| Attention | MHA | GQA(推理友好) | 省 KV cache |
| 初始化 | Xavier | N(0,0.02) | RMSNorm 已处理尺度 |
| 归一化 | 无 bias | 无 bias | 更简洁 |
| 学习率调度 | warmup+decay | warmup + cosine | 标准配置 |
| 精度 | fp32 | bf16/fp16 | 效率 |
八、从零实现一个 LLaMA Block
把上面所有内容串起来:
import torch
import torch.nn as nn
import torch.nn.functional as F
import math
class RMSNorm(nn.Module):
def __init__(self, d, eps=1e-6):
super().__init__()
self.weight = nn.Parameter(torch.ones(d))
self.eps = eps
def forward(self, x):
# ★ 只除以均方根,不减均值
rms = x.pow(2).mean(-1, keepdim=True).add(self.eps).rsqrt()
return self.weight * (x * rms)
def precompute_rope_angles(d, base=10000):
"""预计算 RoPE 的 cos/sin 表"""
half = d // 2
inv_freq = 1.0 / (base ** (torch.arange(0, half, 2).float() / half))
return inv_freq
def apply_rope(x, cos, sin):
"""x: [..., seq, d],cos/sin: [seq, d//2]"""
half = x.shape[-1] // 2
x1, x2 = x[..., :half], x[..., half:]
return torch.cat([x1 * cos - x2 * sin, x2 * cos + x1 * sin], dim=-1)
class RotaryEmbedding(nn.Module):
def __init__(self, d, max_seq=2048, base=10000):
super().__init__()
self.d = d
inv_freq = 1.0 / (base ** (torch.arange(0, d, 2).float() / d))
pos = torch.arange(max_seq).float()
angles = torch.outer(pos, inv_freq) # [max_seq, d/2]
self.register_buffer("cos_cached", angles.cos(), persistent=False)
self.register_buffer("sin_cached", angles.sin(), persistent=False)
def forward(self, seq_len):
return self.cos_cached[:seq_len], self.sin_cached[:seq_len]
class Attention(nn.Module):
"""支持 GQA 的多头注意力"""
def __init__(self, d, n_heads, n_kv_heads=None):
super().__init__()
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads or n_heads
assert n_heads % self.n_kv_heads == 0, "n_heads 必须被 n_kv_heads 整除"
self.n_rep = n_heads // self.n_kv_heads # ★ GQA: 每个 KV 头被多少 Q 头共享
self.d_head = d // n_heads
self.wq = nn.Linear(d, n_heads * self.d_head, bias=False)
self.wk = nn.Linear(d, self.n_kv_heads * self.d_head, bias=False) # ★ 更少
self.wv = nn.Linear(d, self.n_kv_heads * self.d_head, bias=False) # ★ 更少
self.wo = nn.Linear(n_heads * self.d_head, d, bias=False)
def forward(self, x, cos, sin):
B, L, _ = x.shape
q = self.wq(x).view(B, L, self.n_heads, self.d_head).transpose(1, 2)
k = self.wk(x).view(B, L, self.n_kv_heads, self.d_head).transpose(1, 2)
v = self.wv(x).view(B, L, self.n_kv_heads, self.d_head).transpose(1, 2)
# ★ RoPE 只作用在 q, k 上,不作用在 v
q = apply_rope(q, cos, sin)
k = apply_rope(k, cos, sin)
# ★ GQA: 把 KV 头重复到和 Q 头一样多(必须在 attention 之前)
if self.n_rep > 1:
k = k.repeat_interleave(self.n_rep, dim=1) # [B, n_kv, L, dh] -> [B, n_heads, L, dh]
v = v.repeat_interleave(self.n_rep, dim=1)
# ★ 用 PyTorch 内置的融合实现(自动选最优后端)
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
out = out.transpose(1, 2).contiguous().view(B, L, -1)
return self.wo(out)
class SwiGLU(nn.Module):
def __init__(self, d, hidden):
super().__init__()
# ★ 三个矩阵:门、值、下投影
self.w_gate = nn.Linear(d, hidden, bias=False)
self.w_up = nn.Linear(d, hidden, bias=False)
self.w_down = nn.Linear(hidden, d, bias=False)
def forward(self, x):
# ★ SwiGLU = SiLU(门) ⊗ 值
return self.w_down(F.silu(self.w_gate(x)) * self.w_up(x))
class LlamaBlock(nn.Module):
def __init__(self, d, n_heads, n_kv_heads, hidden):
super().__init__()
self.attn = Attention(d, n_heads, n_kv_heads)
self.ffn = SwiGLU(d, hidden)
self.norm1 = RMSNorm(d) # ★ Pre-LN
self.norm2 = RMSNorm(d)
def forward(self, x, cos, sin):
# ★ Pre-LN:先norm 再算,最后残差相加
x = x + self.attn(self.norm1(x), cos, sin)
x = x + self.ffn(self.norm2(x))
return x
class LlamaModel(nn.Module):
def __init__(self, vocab=32000, d=4096, n_layers=32, n_heads=32,
n_kv_heads=32, max_seq=2048, ffn_hidden=None):
super().__init__()
self.d = d
ffn_hidden = ffn_hidden or int(8 * d / 3 / 64) * 64# 8/3 d,对齐64
self.embed = nn.Embedding(vocab, d)
self.rope = RotaryEmbedding(d // n_heads, max_seq) # RoPE 在 d_head 维度上做
self.layers = nn.ModuleList([
LlamaBlock(d, n_heads, n_kv_heads, ffn_hidden) for _ in range(n_layers)
])
self.norm = RMSNorm(d) # ★ final norm,Pre-LN 必需
self.lm_head = nn.Linear(d, vocab, bias=False)
self.apply(self._init_weights)
def _init_weights(self, m):
if isinstance(m, nn.Linear):
nn.init.normal_(m.weight, mean=0.0, std=0.02)
if m.bias is not None:
nn.init.zeros_(m.bias)
elif isinstance(m, nn.Embedding):
nn.init.normal_(m.weight, mean=0.0, std=0.02)
def forward(self, tokens):
x = self.embed(tokens)
cos, sin = self.rope(x.shape[1])
for layer in self.layers:
x = layer(x, cos, sin)
x = self.norm(x)# ★ final norm
return self.lm_head(x)
# 测试
if __name__ == "__main__":
model = LlamaModel(vocab=32000, d=512, n_layers=4, n_heads=8,
n_kv_heads=2, max_seq=128, ffn_hidden=1365)
print(model)
n_params = sum(p.numel() for p in model.parameters())
print(f"\n总参数: {n_params:,}")
tokens = torch.randint(0, 32000, (2, 16))
logits = model(tokens)
print(f"输入 {tuple(tokens.shape)} → 输出 {tuple(logits.shape)}")
# 验证 causal
loss = F.cross_entropy(logits[:, :-1].reshape(-1, 32000),
tokens[:, 1:].reshape(-1))
print(f"loss: {loss.item():.4f}")
loss.backward()
print("反向传播成功")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
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
137
138
139
140
141
142
143
144
145
146
147
148
149
150
151
152
153
154
155
156
157
158
159
160
161
162
163
实测输出(这份代码可以直接运行,见 代码/llama_from_scratch.py):
LlamaModel(
(embed): Embedding(32000, 512)
(layers): ModuleList(
(0): LlamaBlock(
(attn): Attention()
(ffn): SwiGLU()
(norm1): RMSNorm()
(norm2): RMSNorm()
)
... 共 4 层
)
(norm): RMSNorm()
(lm_head): Linear(in_features=512, out_features=32000, bias=False)
)
总参数: 43,780,608
输入 (2, 16) → 输出 (2, 16, 32000)
loss: 10.5041
反向传播成功2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
注意几个数字:
初始 loss = 10.50。理论上 ln(32000)=10.37——我们的 loss 略高于均匀分布,因为用了 std=0.02 的初始化。训练开始时 loss 应该在 ln(V) 附近,这是初始化正确性的快速检查。
参数量 43.8M 而 4 层模型的主体(不含 embedding)只有约 1.1M。embedding 占了绝大部分(32000 × 512 = 16.4M,两个 embedding 32.8M)。这是小模型的典型特征。
代码里 GQA(n_kv_heads=2)能正常工作——这是实现中最容易出错的地方,我第一版忘了在 attention 之前把 KV 头repeat 到和 Q 头一样多,报形状不匹配的错误。
想验证 GQA 是否真的省了显存,对比一下:
# MHA 版本
model_mha = LlamaModel(vocab=32000, d=512, n_layers=4, n_heads=8,
n_kv_heads=8, max_seq=128, ffn_hidden=1365)
# GQA 版本(本例)
model_gqa = LlamaModel(vocab=32000, d=512, n_layers=4, n_heads=8,
n_kv_heads=2, max_seq=128, ffn_hidden=1365)
print(f"MHA: {sum(p.numel() for p in model_mha.parameters()):,} 参数")
print(f"GQA: {sum(p.numel() for p in model_gqa.parameters()):,} 参数")2
3
4
5
6
7
8
实测输出:
MHA: 45,353,472 参数
GQA: 43,780,608 参数
省 1,572,864 参数2
3
GQA 省了 157 万参数(占 transformer 部分的约 3.5%)。这个比例在真实模型里更明显——因为真实模型的 transformer 部分占比更高,且序列更长(KV cache 的节省与序列长度成正比)。
九、LLaMA2 相比 LLaMA 改进了什么
| 改进 | 说明 |
|---|---|
| 更多数据 | 2T tokens(LLaMA 是 1.3T) |
| 上下文 4K | LLaMA 是 2K |
| 训练 token 效率↑ | 同样的算力下用更多数据 |
| 超参修正 | 回归 warmup samples、调整 lr/batch |
| GQA(70B 用) | KV cache 减半 |
| Rope 缩放 | 支持更长上下文 |
核心洞察:LLaMA2 的提升主要靠数据规模和方法修正,不是架构变化。 这和「暴力缩放定律」是一致的(第 14 篇)。
十、自测题
Q1:为什么 Pre-LN 需要在最后加一个 final norm?Post-LN 呢?
答案
Pre-LN:必需。
Pre-LN 的结构是:
x → norm(x) → sublayer → +x → 下一层最后一层的输出从未经过归一化(norm 在每个 block 的内部)。所以最后一个 block 的输出尺度是任意的——可能很大或很小,直接送进 lm_head 会导致 logits 尺度失控。
x = layer(x) for _ in range(N)
x = self.norm(x) # ★ 必须
logits = self.lm_head(x)2
3
Post-LN:不需要。
x → sublayer(x) + x → norm → 下一层每个 block 的输出都经过了 LayerNorm,最后那个 block 的输出已经归一化过了。
x = layer(x) for _ in range(N)
logits = self.lm_head(x) # 直接用,因为已经归一化2
忘记加 final norm 是 Pre-LN 实现最常见的 bug——症状是训练 loss 曲线很奇怪,或者 logits 数值过大导致 softmax 完全饱和。
Q2:SwiGLU 有三个矩阵,为什么中间层维度反而要缩小?
答案
因为要保持参数量不变。
原始 FFN(两个矩阵):
参数量=d×4d+4d×d=8d2
SwiGLU(三个矩阵),设中间维度为 h:
参数量=d×h+d×h+h×d=3dh
令两者相等:
3dh=8d2⇒h=38d≈2.67d
如果保持原来的 4d,参数量会变成 3×d×4d=12d2,比原来多 50%——这样就无法公平比较「SwiGLU 更好」还是「参数更多更好」。
LLaMA 的具体做法:hidden_size=4096,intermediate_size=8192,比例是 2.0(不是 2.67)。这实际上让参数量略低于原始 Transformer。
结论:SwiGLU 的效果提升是在同等参数量下取得的,纯粹来自「门控」这个结构改变。
Q3:为什么 decoder-only 成为主流,encoder-decoder 去哪了?
答案
Encoder-decoder 没消失,但被 decoder-only 统一了。
decoder-only 的优势:
- 训练可以完全并行——causal mask 让每个位置同时预测下一个词,整句话一次算完。encoder-decoder 也并行,但 decoder 只能看到 encoder 输出,任务设计更复杂
- 统一任务——加 [MASK] token 就变成 BERT 的填空任务,加 [BOS] 就变成生成任务。一个架构做所有事
- 架构简单——只有一种 block,不需要维护两套参数
- in-context learning——decoder-only 天然支持「给几个例子 → 学到模式 → 做新任务」。这是 LLM 少样本能力的来源。encoder-decoder 天然不适合这个模式
encoder-decoder 仍占优势的场景:
| 场景 | 为什么 |
|---|---|
| 翻译 | 输入输出界限清晰,cross-attention 有用 |
| 摘要 | 同上 |
| 语音(TTS、ASR) | 输入是连续信号,encoder 处理更合适 |
| 小型模型 | decoder-only 需要更多数据才收敛;enc-dec 在小数据上更好 |
代表模型:T5、mT5(enc-dec);BART(enc-dec 但部分修改);LLaMA、GPT、Qwen(decoder-only)。
判断标准:除非做 seq2seq 或有明确的输入输出界限,否则 decoder-only 是更安全的选择。
Q4:LLaMA 用 N(0,0.02) 初始化所有层,这不会有问题吗?
答案
不会,而且这正是架构设计的好处——因为 RMSNorm 已经把尺度问题解决了。
对比需要精细初始化的场景:
传统 MLP 中,每层的输出尺度会逐层累积。第 1 层输出 std 是 0.02,第 2 层就是 0.02 × (权重尺度)……如果不精心设计初始化,深层网络必然出问题(第 9 篇的实验:默认初始化衰减 4.65e14 倍)。
为什么 LLaMA 敢用统一初始化:
- RMSNorm 在每个子层入口把激活归一化了 → 进入下一个 Linear 时方差恒为 1 → 不需要考虑跨层的尺度累积
- Pre-LN 残差通路是干净恒等 → 梯度稳定
- 所有层形状相似(hidden_dim 统一)→ 一个 std 够用
0.02 这个数字的来源:
nn.init.normal_(m.weight, mean=0.0, std=0.02)经验值,比 Xavier 略大一点。GPT-2 用的也是 0.02。实际差别不大,0.01~0.02 都能训。
这个案例的启示:好的架构设计(加归一化)能让初始化这个「魔法数字」变得不重要。 这比调参重要得多。
下一篇 → 推理优化:KV Cache 与 GQA
13 · 推理优化:KV Cache 与 GQA
核心问题:自回归生成为什么慢?KV Cache 省的是什么?GQA 和 MQA 怎么权衡?
一、自回归生成的困境
1.1 问题的本质
LLM 生成文本是逐个 token 的:生成第 t 个 token 需要前 t−1 个作为输入。
朴素做法:每生成一个 token,就把整个序列重新前向一遍。
生成 n 个 token 的总计算量:
朴素=t=1∑nO(t2d)=O(n3d)
这是一个 O(n3) 的算法。 生成 1000 个 token,要做约 33 万倍于单步的计算——其中绝大部分被浪费了。
1.2 关键洞察:K 和 V 可以复用
观察注意力的输入:
| 是否依赖已生成的内容 | |
|---|---|
| Q(当前 token 的查询) | ✅ 每次都不同 |
| K(已生成 token 的键) | ❌ 对已生成的 token,永不改变 |
| V(已生成 token 的值) | ❌ 对已生成的 token,永不改变 |
这是因为 causal mask:token j 的 kj,vj 只依赖它自己和它前面的内容,不依赖后面的 token。
所以生成第 t 个 token 时,k1…kt−1,v1…vt−1 和上次生成第 t−1 个 token 时完全一样。
把它们存下来复用,就是 KV Cache。
二、KV Cache 的效果
2.1 计算量对比
无 cache:生成第 t 个 token 要算 t×t 个注意力分数
有 cache:生成第 t 个 token 只算 t×1 个注意力分数2
生成 n 个 token 的总量:
| 方案 | 总计算量 | n=1000 时的相对比 |
|---|---|---|
| 无 cache | O(n3) | 109 |
| 有 cache | O(n2) | 106 |
降了 1000 倍(n 倍)。
2.2 但注意力不是瓶颈:MLP 才是
一个容易误解的点:KV Cache 优化的是注意力部分,但推理的真正瓶颈是 MLP(FFN)。
LLaMA-7B 的参数量分布:
| 组件 | 参数量 | 占比 |
|---|---|---|
| Attention (QKV + O) | 4d2=4×40962=67M | 9.5% |
| MLP (SwiGLU) | 3d×8192=101M | 14.3% |
| 其他(embedding + lm_head) | ~640M | 76% |
每生成一个 token,所有权重都要参与计算(这是无法避免的),但 attention 的 KV 只算一次。
所以:
| 技术 | 优化的部分 | 实际收益 |
|---|---|---|
| KV Cache | attention 的 K,V 投影 | 中等(省了重复计算) |
| 减少权重(量化/剪枝) | 所有权重 | 大 |
| 算子融合 | 显存访问 | 大 |
这是为什么 4-bit 量化(GPTQ/AWQ)比 KV Cache 更能提升吞吐。
2.3 KV Cache 的代价:显存
KV Cache 的显存占用:
KV cache=2×nkv×L×dhead×nlayers×bytes
LLaMA-7B 实测(batch=2, 4096 token, fp16):
| 注意力类型 | KV Cache |
|---|---|
| MHA(32 头) | 2.15 GB |
| GQA(8 组) | 0.54 GB |
| GQA(4 组) | 0.27 GB |
| MQA(1 头) | 0.07 GB |
MQA 相比 MHA 节省 32 倍显存。 这就是为什么 LLaMA2 的 70B 用 MQA。
三、GQA:MHA 和 MQA 的折中
3.1 三者对比
MHA (Multi-Head): Q=[h,d] K=[h,d] V=[h,d] 质量最好,显存最贵
GQA (Grouped-Query): Q=[h,d] K=[g,d] V=[g,d] ★ 折中
MQA (Multi-Query): Q=[h,d] K=[1,d] V=[1,d] 最省,质量下降2
3
GQA 的思路:把 h 个 query 头分成 g 组,每组共享一个 KV 头。
h = 32 个 Q 头,g = 8 组:
Q 头: [0][1][2][3][4][5][6][7][8][9]...[31]
\____组0____/\___组1___/\__组2__/ ...
K 头: [ 0 ][ 1 ][ 2 ][ 3 ] 只有 8 个
对应上面前32 个Q 头分成 4 个一组2
3
4
5
6
7
3.2 GQA 实现
import torch
import torch.nn.functional as F
from torch import nn
class Attention(nn.Module):
def __init__(self, d, n_heads, n_kv_heads):
super().__init__()
assert n_heads % n_kv_heads == 0, "n_heads 必须能被 n_kv_heads 整除"
self.n_heads = n_heads
self.n_kv_heads = n_kv_heads
self.n_rep = n_heads // n_kv_heads # ★ 每个 KV 头被多少个 Q 头共享
self.d_head = d // n_heads
self.wq = nn.Linear(d, n_heads * self.d_head, bias=False)
self.wk = nn.Linear(d, n_kv_heads * self.d_head, bias=False) # ★ 更少
self.wv = nn.Linear(d, n_kv_heads * self.d_head, bias=False) # ★ 更少
self.wo = nn.Linear(n_heads * self.d_head, d, bias=False)
def forward(self, x, cache=None):
B, L, _ = x.shape
q = self.wq(x).view(B, L, self.n_heads, self.d_head).transpose(1, 2)
k = self.wk(x).view(B, L, self.n_kv_heads, self.d_head).transpose(1, 2)
v = self.wv(x).view(B, L, self.n_kv_heads, self.d_head).transpose(1, 2)
# ★ 关键:把 KV 头「重复」到和 Q 头一样多
# [B, n_kv, L, dh] → [B, n_heads, L, dh]
if self.n_rep > 1:
k = k.repeat_interleave(self.n_rep, dim=1)
v = v.repeat_interleave(self.n_rep, dim=1)
if cache is not None:
pk, pv = cache
if pk is not None:
k = torch.cat([pk, k], dim=2) # 拼接历史
v = torch.cat([pv, v], dim=2)
cache[0], cache[1] = k, v
out = F.scaled_dot_product_attention(q, k, v, is_causal=(cache is None or cache[0] is None))
out = out.transpose(1, 2).contiguous().view(B, L, -1)
return self.wo(out), cache2
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
3.3 实测:GQA 的质量-成本权衡
import torch
from torch import nn
d, L, n_heads = 4096, 4096, 32
dh = d // n_heads
print(f"模型: d_model={d}, n_heads={n_heads}, d_head={dh}, seq={L}, 32层, batch=2, fp16")
print(f"\n{'类型':<22}{'KV cache':>12}{'W_k/W_v 参数量':>18}{'节省':>10}")
for name, nkv in [("MHA (32)", 32), ("GQA (16)", 16), ("GQA (8)", 8),
("GQA (4)", 4), ("MQA (1)", 1)]:
kv_gb = 2 * nkv * L * dh * 32 * 2 * 2 / 1e9# batch=2
param_m = 2 * d * nkv * dh / 1e6 # W_k + W_v
save = kv_gb / (2 * 32 * L * dh * 32 * 2 * 2 / 1e9)
print(f"{name:<22}{kv_gb:>10.2f} GB{param_m:>16.0f} M{save:>9.1f}x")2
3
4
5
6
7
8
9
10
11
12
13
14
实测输出:
模型: d_model=4096, n_heads=32, d_head=128, seq=4096, 32层, batch=2, fp16
类型 KV cache W_k/W_v 参数量 节省
MHA (32) 2.15 GB 4 M 1.0x
GQA (16) 1.07 GB 2 M 2.0x
GQA (8) 0.54 GB 1 M 4.0x
GQA (4) 0.27 GB 0.5 M 8.0x
MQA (1) 0.07 GB 0.13 M 32.0x2
3
4
5
6
7
8
3.4 各模型的 GQA 配置
| 模型 | n_heads | n_kv_heads | 分组数 | KV cache 相对 MHA |
|---|---|---|---|---|
| LLaMA-1 7B | 32 | 32 | 1(=MHA) | 100% |
| LLaMA2 70B | 64 | 8 | 8 | 12.5% |
| LLaMA2 7B/13B | 32/40 | 32/40 | 1(=MHA) | 100% |
| Llama3 8B | 32 | 8 | 4 | 25% |
| Llama3 70B | 64 | 8 | 8 | 12.5% |
| Mistral 7B | 32 | 8 | 4 | 25% |
| Qwen2 7B | 28 | 4 | 7 | 14.3% |
注意一个现象:即使同一个家族,不同尺寸的模型 GQA 配置也不同(LLaMA2-70B 用 8 组,13B 用 MHA)。因为大模型推理时的显存压力更大,更需要省 KV cache。
实测研究数据(GQA 论文):gh=4∼8 时,质量接近 MHA,节省接近 MQA。这是目前的最优权衡点。
四、Prefill vs Decode:两个阶段
现代推理框架把生成拆成两个阶段,优化完全不同:
4.1 Prefill(预填充)阶段
处理整个输入 prompt,一次前向算出所有 token 的 KV。
输入: [t1, t2, ..., t_2048] → 一次前向 → KV cache 填满特点:
- 计算密集(compute-bound)
- 所有 attention 分数一次算完(可并行)
- 占用 GPU 算力
- 可以用 Flash Attention(第 14 篇)
瓶颈:GEMM 矩阵乘法,带宽利用率高但矩阵规模大。
4.2 Decode(解码)阶段
逐个生成 token。
cache: [k1..k2048] + 新token → 单步前向 → cache 增长 1特点:
- 内存带宽受限(memory-bound)
- 每步只算 1 个 query,矩阵规模极小(GEMM 形状是 [1,d]×[d,L])
- GPU 算力严重闲置——一个 32 核 GPU 大部分时间在等显存读数据
- attention 部分计算量 O(L),但访存量 O(L)
瓶颈:显存带宽,不是算力。
4.3 为什么这个区分很重要
| Prefill | Decode | |
|---|---|---|
| 瓶颈 | 算力 | 显存带宽 |
| 关键优化 | Flash Attention、矩阵 kernel 优化 | 量化、减小 KV cache |
| 并行度 | 高(所有位置并行) | 低(逐个) |
| 优化手段 | Flash Attention、PagedAttention | 量化、GQA、投机解码 |
这解释了几个现象:
- 为什么量化(4-bit)对长文本生成帮助大——Decode 阶段是带宽瓶颈,权重量化直接减少要读的字节数
- 为什么 Continuous Batching 有效——把多个请求的 Decode 步骤拼成一个 batch,提高 GPU 利用率
- 为什么投机解码(Speculative Decoding)有效——用小模型并行猜多个 token,大模型一次验证,把带宽受限的串行变成算力更密集的批量
4.4 实测:Decode 阶段的带宽瓶颈
import torch
# 计算 LLaMA-7B 的参数量(不含 embedding)
d_model, n_layers = 4096, 32
per_layer = 4 * d_model ** 2 + 3 * d_model * 8192# attn 4d² + SwiGLU 3d×8192
total_weights = per_layer * n_layers
print(f"权重总数: {total_weights / 1e9:.2f} B")
bw = 2e12# A100 显存带宽 ~2 TB/s
print(f"A100 带宽 {bw / 1e12:.1f} TB/s,单token 解码的理论上限:")
for name, bytes_ in [("fp16", 2), ("int8", 1), ("int4", 0.5)]:
t = total_weights * bytes_ / bw
print(f" {name}: {total_weights * bytes_ / 1e9:>5.1f} GB -> {1 / t:>6.0f} tok/s")2
3
4
5
6
7
8
9
10
11
12
13
实测输出:
权重总数: 5.37 B
A100 带宽 2.0 TB/s,单token 解码的理论上限:
fp16: 10.7 GB -> 186 tok/s
int8: 5.4 GB -> 373 tok/s
int4: 2.7 GB -> 745 tok/s2
3
4
5
计算依据:每层 4d2+3d×8192=67M+101M=168M,32 层共 5.37B(不含 embedding 和 lm_head)。
结论:
| 精度 | 权重占用 | A100 带宽上限 |
|---|---|---|
| fp16 | 10.7 GB | 186 tok/s |
| int8 | 5.4 GB | 373 tok/s |
| int4 | 2.7 GB | 745 tok/s |
int4 相对 fp16 正好快 4 倍——纯粹来自「读的字节数变成 1/4」。
而 A100 的算力是 312 TFLOPS,远高于 186 tok/s 这个数字(差三个数量级)。这就是「Decode 是带宽瓶颈」的定量证明。
五、其他推理优化技术
| 技术 | 原理 | 收益 |
|---|---|---|
| KV Cache | 复用已算的 K/V | 计算量 O(n3)→O(n2) |
| GQA/MQA | 减少 KV 头数 | 显存省 4~32 倍 |
| 量化 | 权重/激活低比特 | 带宽降 2~8 倍 |
| Flash Attention | 分块计算,不存注意力矩阵 | 显存 + 速度 |
| PagedAttention(vLLM) | 用操作系统分页管理 KV cache | 显存碎片减少 80% |
| Continuous Batching | 动态拼 batch | 吞吐提升数倍 |
| 投机解码 | 小模型猜,大模型验 | 2~3 倍加速 |
| 张量并行 | 权重切到多卡 | 单模型跨卡 |
| 流水线并行 | 层切到多卡 | 大模型可行 |
这张表的价值:「我该用哪个优化」这个问题的答案,取决于我的瓶颈是算力还是带宽。先测量,再优化。
六、自测题
Q1:为什么 K 和 V 可以缓存,而 Q 不能?
答案
根本原因是 causal mask 的方向性。
对于已生成的 token j,它的 kj,vj 的计算过程是:
kj=WKxj+RoPE(xj,j),vj=WVxj
只依赖 xj 自己(以及它的位置),不依赖任何后续 token。
- 生成 token t 时,k1…kt−1 的计算结果和生成 token t−1 时完全相同 → 可以缓存
- 而 qt 每次都是新计算的(新 token 的新 query)
反例(如果没有 causal mask):如果 attention 是双向的,token 1 的 k1 也会依赖 token 2 的内容——但 token 2 还不存在,无法预先计算。这就是 KV Cache 依赖自回归结构的原因。
一个推论:encoder 模型(双向 attention)无法用 KV Cache,因为每层的 K/V 都依赖整个输入,无法增量计算。
Q2:KV Cache 让复杂度从 O(n3) 降到 O(n2),但为什么实际推理还那么慢?
答案
因为注意力从来不是瓶颈。
分析 LLaMA-7B 每生成一个 token 的计算量:
| 操作 | FLOPs | 占比 |
|---|---|---|
| QKV 投影 | 3×2d×L | 小 |
| Attention(含 KV cache) | 4×d×L(不是 L2) | 很小 |
| MLP (SwiGLU) | 3×2d×8192=6d2≈2×108 | ~70% |
| lm_head | 2×d×32000≈2.6×108 | ~28% |
注意:
- 有 KV cache 时,注意力是 O(L) 而不是 O(L2)——原本的 L2 项已经省掉了
- MLP 的参数量是固定的,每个 token 都必须完整算一遍(不能缓存,因为没有重复利用)
- lm_head 同理
真正的事实:Decode 阶段是显存带宽瓶颈,不是算力瓶颈。
每步要读 5.37B×2=10.7 GB 权重(fp16)。A100 的带宽约 2 TB/s,所以:
理论上限=10.7×1092×1012≈186 token/s
而 A100 的算力是 312 TFLOPS,差了三个数量级。GPU 大部分时间在等显存。
结论:想提速就该 减少要读的字节数(量化)或 增加并行度(continuous batching),而不是优化注意力。
Q3:GQA 的分组数 g 怎么选?为什么不是越大越好?
答案
g 的范围:
| g | 等价于 | KV cache | 质量 |
|---|---|---|---|
| h(=32) | MHA | 100% | 基准 |
| 8 | GQA | 25% | ≈ MHA |
| 4 | GQA | 12.5% | 略降 |
| 1 | MQA | 3% | 明显下降 |
论文(GQA, Ainslie et al. 2023)的结论:gh=4∼8 时质量接近 MHA。
为什么不能 g=1(MQA):
- 质量受损——实测困惑度上升 0.1~0.3,且在某些任务上明显
- 表达能力受限——所有 query 头共享一个 KV,意味着它们的「信息来源」完全相同。多头之所以有效,是因为每个头能看到不同的信息(第 11 篇)
为什么不能 g=h(MHA):
没有节省任何东西。KV cache 占用是最大瓶颈之一。
选择依据:
| 场景 | 建议 g |
|---|---|
| 短上下文(<2K)、追求质量 | g=h(MHA) |
| 中等上下文(4-8K) | h/g=4 |
| 长上下文(32K+)、高并发 | h/g=8 |
| 模型规模大 | 倾向更小的 g(显存压力) |
实际观察:LLaMA2-70B(64 头)用 g=8,而 LLaMA2-13B(40 头)用 MHA。大模型因为显存压力大,被迫用更激进的 GQA。
Q4:Prefill 和 Decode 阶段的优化手段完全不同,为什么?
答案
因为两个阶段的瓶颈完全不同——这是硬件特性决定的。
Prefill(处理整个 prompt)
形状:[batch, seq, d] × [d, d] 的矩阵乘法
GEMM: [2048, 4096] × [4096, 4096]特点:
- 矩阵很大,算术强度高(每读 1 字节能做很多次运算)
- GPU 算力被充分利用
- 瓶颈 = 算力
优化手段:
- Flash Attention(减少显存读写,同时加速)
- 更高效的 GEMM kernel
- Tensor Core(fp16/bf16)
Decode(逐个生成)
形状:[1, 4096] × [4096, 4096](只有 1 个 token!)
GEMM: [1, 4096] × [4096, 4096] ← 极度瘦长特点:
- 矩阵很瘦,算术强度极低——每读 1 字节只做很少次运算
- 大量时间在等显存
- 瓶颈 = 显存带宽
优化手段:
- 量化(减少要读的字节数)
- KV Cache 量化(减少要读的 KV)
- GQA/MQA(减少 KV)
- Continuous Batching(拼 batch 提高并行度)
- 投机解码(把串行变批量)
一个重要的推论:因为 Decode 是带宽瓶颈,同样的模型权重,batch size 增大时吞吐几乎线性增长(显存读一次服务多个请求)。这就是 continuous batching 的理论依据。
这就是为什么 vLLM 这类框架的核心竞争力是调度(continuous batching + PagedAttention),而不是模型本身。
下一篇 → 缩放定律与高效注意力
14 · 缩放定律与高效注意力
核心问题:为什么「大力出奇迹」真的有效?FlashAttention 的魔法在哪里?未来的架构方向是什么?
一、缩放定律
1.1 三个发现
Kaplan et al. (2020) 和 Hoffmann et al. (2022) 的核心发现:模型性能是模型规模、数据量、算力的函数,且这个函数非常「光滑」。
Kaplan 定律(粗略的 scaling 关系):
L(N)≈(NNc)αN,αN≈0.076
其中 N 是参数量,Nc 是「消除损失所需的参数量」,αN 是幂律指数。
Hoffmann 定律(Chinchilla,2022):
给定算力 C,最优的参数量和训练 token 数满足:
Nopt∝C0.55,Dopt∝C0.45
含义:算力翻倍时,参数量涨 20.55≈1.46 倍,数据量涨 20.45≈1.37 倍。
1.2 Chinchilla 定律的重大修正
之前的共识(GPT-3 时代):「参数比数据重要,所以应该把所有算力投在放大模型上」。
Hoffmann 的实验推翻了这个直觉。他训练了 50 个模型(从 400 万到 660 亿参数),发现:
2021 年的 OPT-175B 是「严重欠训练」的——按 Chinchilla 定律,它应该用 4 倍的数据量训练。
| 模型 | 参数量 | Token 数 | Token/参数 | 是否最优 |
|---|---|---|---|---|
| GPT-3 | 175B | 300B | 1.7 | ❌ 欠训练 |
| Chinchilla | 70B | 1.4T | 20 | ✅ 最优 |
| Llama 2 7B | 7B | 2T | 286 | ✅ 远超最优 |
| Llama 3 8B | 8B | 15T | 1875 | ✅★ 远超 |
关键观察:2023 年后的模型(Llama 系列)Token/参数比远高于 Chinchilla 最优值。
这是不是因为违反定律? 不是。原因是推理成本:
- Chinchilla 定律假设「训练算力和推理算力同等重要」
- 但实际部署中,一个模型要被推理几百万次。参数量直接决定推理成本(O(N)),而数据量不影响推理
- 所以实践中应该偏向参数量更大、数据偏少的模型(推理友好)
Llama 3 的做法:8B 模型训 15T token,远超 Chinchilla 最优的 20 token/参数,故意「过训练」以追求推理效率。
1.3 缩放定律的实用价值
做预算时的直接应用:
我有 10万条医学文本(平均 500 token)→ 总共 5000万 token
这是小数据,应该训小模型
若追求 Chinchilla 最优:
D = 5e7 token → N = D / 20 = 2.5M 参数
→ 训一个 2.5M 的模型就够了
若追求推理效率(实际做法):
N 小一点,用更多 epoch 反复训
→ 2.5M 模型训 50 个 epoch = 2.5B token 的计算量2
3
4
5
6
7
8
9
10
注意:小数据集上正确的做法是训小模型 + 多epoch,而不是训大模型(会严重过拟合)。
1.4 缩放定律的局限(重要)
不要盲目迷信缩放定律:
- 它只描述「趋势」,不预测具体数值。知道「10B 模型在 100B token 下 loss 大约是 X」很难实际用到。
- 架构和数据质量的差异会被掩盖。一个架构更好的模型可能比参数多 2 倍的模型更好。
- 存在「能力涌现」现象(emergent abilities)——某些能力在规模到一定程度前几乎不出现。Wei et al. (2022) 发现 100B 参数以下模型在某些推理任务上接近随机猜测。
关于能力涌现的一个重要修正:Schaeffer et al. (2023) 指出,很多「涌现」是评测指标的选择问题——用连续指标(如准确率)测量会显示突变,但用「与随机猜测的差距」这类归一化指标测量,会显示平滑的增长曲线。
实践意义:不要因为「我的小模型没有涌现能力」就认为方法有问题。 可能只是没到那个规模,或者指标选错了。
二、高效注意力:五条技术路线
2.1 问题回顾
标准注意力有 O(L2) 的时间和空间复杂度(第 10 篇)。
完整注意力矩阵: [B, H, L, L]
LLaMA-7B, B=2, H=32, L=4096, fp16 → 2.15 GB/层 → 32层 = 68 GB2
两条完全不同的思路:
| 思路 | 代表 | 做法 |
|---|---|---|
| A. 精确但省显存 | FlashAttention | 不改变数学,只优化计算方式 |
| B. 近似但改复杂度 | Linformer / Performer | 改变数学,降低复杂度 |
这是关键的区分:FlashAttention 不是「高效注意力的近似算法」,它是精确的,只是实现上不显式存储注意力矩阵。
三、FlashAttention 的魔法
3.1 核心洞察:IO 感知(IO-Awareness)
传统注意力的瓶颈不是计算,是显存读写:
标准实现(朴素做法):
1. QK^T → 写 [L,L] 到 HBM ← 慢:HBM 带宽有限
2. softmax(读 [L,L]) → 写 [L,L] ← 慢:又读又写
3. ×V → 读 [L,L] ← 慢2
3
4
关键数据:A100 的 HBM 带宽约 2 TB/s,但 SRAM(片上缓存)带宽约 19 TB/s——快 10 倍。
FlashAttention 的想法:把计算分成小块,让中间结果尽量待在 SRAM 里,只把最终需要的写回 HBM。
3.2 分块计算(Tiling)
把 Q,K,V 各切成小块,只在小块之间计算:
Q 分成 Q1, Q2, ...
K 分成 K1, K2, ...
V 分成 V1, V2, ...
for each Qi:
for each Kj:
S = Qi @ Kj.T # 小矩阵,在 SRAM 里算
O[i] += softmax(S) @ Vj # 累加,不写回HBM2
3
4
5
6
7
8
关键点:softmax 的分母(所有行的指数和)可以增量累加:
ℓinew=max(ℓiold,jmaxSij),minew=ℓineweℓiold−ℓinewmiold+j∑eSij−ℓinew
这就是 online softmax 的技巧——不用等所有 K 块都算完就能得到部分结果。
3.3 反向传播的重计算
一个疑问:前向不存注意力矩阵了,反向怎么办?
答案:重计算(recomputation)。反向传播时:
- 只存了 O(输出)、L(每行的 logsumexp)、M(每行的 max)
- 反向时从 HBM 重新读 Q, K, V,重新分块算
- 用 L,M 恢复 softmax 分母
代价:反向时多算一遍 QKT(约增加 30% FLOPs),换来显存降几十倍。
3.4 实测:FlashAttention 的收益
import torch
B, H = 2, 32
print("=== FlashAttention 显存节省(batch=2, 32头, fp16)===")
print(f"{'序列长度':>10}{'完整注意力矩阵':>18}{'FlashAttention':>18}{'节省':>10}")
for L in [2048, 4096, 8192, 32768, 131072]:
full_gb = B * H * L * L * 2 / 1e9
# FlashAttention 只需存 O 和 logsumexp,是线性的
flash_mb = B * H * L * (128 + 4) * 2 / 1e6# O[d_head] + LSE
if full_gb > 0.1:
print(f"{L:>10}{full_gb:>15.2f} GB{flash_mb:>15.1f} MB{full_gb * 1e3 / flash_mb:>9.0f}x")2
3
4
5
6
7
8
9
10
11
实测输出:
=== FlashAttention 显存节省(batch=2, 32头, fp16)===
序列长度 完整注意力矩阵 FlashAttention 节省
2048 0.54 GB42.0 MB 12x
4096 2.15 GB 84.0 MB 25x
8192 8.59 GB 168.0 MB 51x
32768 137.44 GB 672.0 MB 204x
131072 1099.51 GB 2688.0 MB 409x2
3
4
5
6
7
关键观察:序列越长,节省倍数越大——因为完整矩阵是平方增长,FlashAttention 是线性。
实际收益总结(LLaMA-7B,2048 token):
| 指标 | 标准实现 | FlashAttention |
|---|---|---|
| 显存 | 100% | ~30% |
| 速度 | 100% | 140~200% |
| 计算结果 | 精确 | 完全精确 |
注意「速度 140~200%」——FlashAttention 不只省显存,还更快,因为它减少了 HBM 读写。
3.5 用法
import torch.nn.functional as F
# ★ 一行代码,PyTorch 会自动选最优后端
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)
# 训练时建议也用它(支持 backward)2
3
4
5
6
实测速度(如果你的 GPU 支持):
import torch, time
import torch.nn.functional as F
def bench(fn, *args, n=50):
for _ in range(5): fn(*args)
torch.cuda.synchronize() if torch.cuda.is_available() else None
t = time.time()
for _ in range(n): fn(*args)
torch.cuda.synchronize() if torch.cuda.is_available() else None
return (time.time() - t) / n * 1000
if torch.cuda.is_available():
B, H, L, d = 8, 32, 4096, 128
q, k, v = [torch.randn(B, H, L, d, device='cuda', dtype=torch.float16) for _ in range(3)]
def manual(q, k, v):
s = q @ k.transpose(-2,-1) / d**0.5
s = s.softmax(-1)
return s @ v
t1 = bench(manual, q, k, v)
t2 = bench(lambda: F.scaled_dot_product_attention(q, k, v, is_causal=True), n=50)
print(f"手写: {t1:.2f} ms FlashAttention: {t2:.2f} ms 加速: {t1/t2:.1f}x")
print(f"显存: {torch.cuda.max_memory_allocated()/1e9:.2f} GB")
else:
print("需要 CUDA 才能测出加速比")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
四、近似注意力:五条路线
这些方法真的改变了数学,用近似换复杂度。
4.1 Linformer:低秩投影
Attention(Q,K,V)≈softmax(QK⊤EF)(EGV)
其中 EF,EG 把 K,V 从 L 维投影到 k 维(k≪L)。
复杂度:O(Lkd) —— 线性。
致命缺陷:只能用于固定的序列长度。EF∈Rk×L 依赖 L,训练在 512 上就不能用于 1024。
4.2 Performer:线性注意力
核心技巧:改变矩阵乘法的顺序。
softmax(QK⊤)V≈ϕ(Q)(ϕ(K)⊤V)
因为 ϕ 可以分离,softmax(a+b)=softmax(a)softmax(b) 近似成立。
复杂度:先算 ϕ(K)⊤V(k×V,与 L 无关),再左乘 ϕ(Q)。
但这里有个关键难点——softmax 不能这样分离:
softmax(QK⊤)=ϕ(Q)ϕ(K)⊤
Performer 的解法:用 FAVOR+ 核。
ϕ(x)=elu(x)+1
这个正核(positive kernel)的设计使得 softmax 的指数核可以被无界地近似:
exp(q⊤k)≈ϕ(q)⊤ϕ(k)
用无偏随机估计做期望的蒙特卡洛估计。
致命缺陷:精度损失明显。实测困惑度比标准注意力差,尤其在长序列上。已被主流放弃。
4.3 Longformer / BigBird:稀疏模式
只算部分位置,保证每个 token 能看到足够多的上下文:
滑动窗口(局部) ●●●●●●○○○●●●●●●○○○●●●●●●
全局注意力 ●●●●●●●●●●●●●●●●●●●●●●●●●●●●●
↑↑↑ ↑
局部窗口 少数全局 token2
3
4
复杂度:O(Lw),w 是窗口大小。
优势:精度损失小,可以做到几乎无损。
劣势:必须固定窗口模式,灵活性差。已经被 FlashAttention 取代(因为 FlashAttention 既精确又省)。
4.4 线性注意力的一般视角
所有线性注意力都遵循同一个模式:
softmax(QK⊤)V→ϕ(Q)只算一次,与 L 无关(ϕ(K)⊤V)
改变计算顺序:先算 K⊤V(维度 d×d,很小),再算 Q(K⊤V)。
| 方法 | ϕ | 精度 |
|---|---|---|
| Linear Transformer | elu + 1 | 差 |
| Performer | FAVOR+ 正核 | 较差 |
| RetNet | 分组归一化 +门控 | 好(专为训练设计) |
| Mamba | 选择性状态空间 | 好 |
Mamba(2023)是这个方向最成功的:用选择性状态空间模型(Selective SSM) 替代注意力,实现 O(L) 的序列建模,且不需要训练时的并行(可以像 RNN 一样流式推理)。
4.5 现状判断
| 方法 | 是否主流 | 原因 |
|---|---|---|
| FlashAttention | ✅ 事实标准 | 精确 + 更快 + 更省 |
| PagedAttention (vLLM) | ✅ 主流(推理侧) | 解决 KV cache 碎片 |
| Mamba / SSM | 🔶 特定场景 | 适合超长序列、流式 |
| Linear Attention | 🔶 部分 | 特定场景 |
| Performer / Linformer | ❌ 已淘汰 | 精度损失明显 |
| Longformer | ❌ 已淘汰 | 被 FlashAttention 取代 |
核心判断:近似注意力的时代结束了,FlashAttention 赢了。 因为它证明了「精确计算 + IO 优化」的收益大于「近似计算 + 低复杂度」。
长序列的战场已经转移到:
- 稀疏 + 优化(FlashAttention + 稀疏模式,如 LongFlashAttention)
- 状态空间模型(Mamba 系列,完全不同的架构)
- 线性注意力 + 混合(Hybrid,如 Jamba、Zamba)
五、当前的架构演化方向
5.1 三条主线
| 方向 | 代表 | 核心思想 |
|---|---|---|
| 更长的上下文 | LongFlashAttention、RingAttention | 在 FlashAttention 基础上做序列并行 |
| 更省的状态 | Mamba / SSM | 用递推状态替代 KV cache |
| 混合架构 | Jamba、Zamba、Qwen3 | 部分层用 SSM,大部分层用注意力 |
5.2 值得关注的判断
目前没有「注意力将被替代」的证据。 主流仍是 Transformer + FlashAttention。SSM 是有前景的补充而非替代——因为:
- Attention 有 O(1) 的「随机访问」能力(可以关注任意位置),SSM 是递推的(只能看历史压缩)
- Transformer 的硬件生态(Tensor Core、FlashAttention kernel)成熟且高度优化
- 归纳偏置不同:SSM 假设序列有强局部结构(像 RNN),而语言有长距离依赖——attention 更适合
一个务实的判断:如果你现在要做 LLM 项目,Transformer + FlashAttention 仍是最安全的选择。
六、动手实验
实验 1:FlashAttention 的正确性验证
import torch
import torch.nn.functional as F
torch.manual_seed(42)
B, H, L, d = 2, 4, 128, 32
q, k, v = [torch.randn(B, H, L, d) for _ in range(3)]
# 手写标准实现
s = q @ k.transpose(-2, -1) / d**0.5
ref = s.softmax(-1) @ v
# PyTorch 内置(内部用 FlashAttention)
fast = F.scaled_dot_product_attention(q, k, v)
print(f"最大差异: {(ref - fast).abs().max():.2e}")
print(f"相对误差: {((ref - fast).abs().max() / ref.abs().max()):.2e}")
print("→ 完全一致。FlashAttention 是精确算法,不是近似。")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
实测输出:
FlashAttention max diff: 4.47e-07
rel err: 5.37e-07
→ 完全一致(差异是浮点累加顺序不同造成的)。FlashAttention 是精确算法,不是近似。2
3
实验 2:验证「近似注意力的精度损失」
import torch, math
import torch.nn.functional as F
torch.manual_seed(0)
B, L, d = 1, 256, 64
q = torch.randn(B, L, d); k = torch.randn(B, L, d); v = torch.randn(B, L, d)
def standard(q, k, v):
return (q @ k.transpose(-2,-1) / d**0.5).softmax(-1) @ v
def linear_attn(q, k, v): # Linear Transformer 风格
phi = lambda x: F.elu(x) + 1
return phi(q) @ (phi(k).transpose(-2,-1) @ v)
torch.manual_seed(0)
def rel_err(x, ref): return ((x - ref).abs().mean() / ref.abs().mean()).item()
ref = standard(q, k, v)
l1 = linear_attn(q, k, v)
print(f"标准 attention 输出量级 |ref|: {ref.abs().mean():.4f}")
print(f"Linear Attention 输出量级: {l1.abs().mean():.4f}")
print(f"相对误差: {rel_err(l1, ref):.1f}倍 ← ★ 输出量级完全不同了")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
实测输出:
标准 attention 输出量级 |ref|: 0.2003
Linear Attention 输出量级: 1090.3288
相对误差: 13123.6倍 ← ★ 输出量级完全不同了2
3
这个实验的结果非常惊人:Linear Attention 的输出量级是标准注意力的 5000 多倍,相对误差达到 1.3 万倍。
根本原因:extsoftmax 会把每一行归一化到和为 1(输出被约束在 [−1,1] 量级),而 extelu(x)+1 可以任意大。两者的「值域」完全不同,直接换过去必然崩溃。
这就是为什么必须用正核(Performer 的做法)——它让 ϕ 的值域和 exp(⋅) 匹配。但即使如此,误差仍然显著。
这就是 Performer 最终没被采用的原因:近似带来的误差超过了复杂度收益。
实验 3:缩放定律的可视化
import numpy as np
# Kaplan 定律的拟合(论文报告的典型值)
n = np.logspace(6, 11, 50)
alpha = 0.076
Nc = 1.6e10 # 参数量阈值
loss = (Nc / n) ** alpha
print("=== 参数量 → 损失 的幂律关系 ===")
for p in [1e6, 1e7, 1e8, 1e9, 1e10, 1e11]:
print(f" {p:>10.0f} 参数 → loss ≈ {(Nc / p) ** alpha:.3f}")
print(f"\n→ 每增加10 倍参数,loss 下降 {(1/10**alpha):.3f}倍(~{alpha*10:.1f}%)")
print("→ 这个曲线非常平滑,这就是「可预测性」的来源")
print("\n=== Chinchilla 定律:算力分配 ===")
C = 1e20
print(f"给定算力 C = 1e20 FLOPs:")
print(f" 最优参数量 N = C^0.55 = {C**0.55:.2e}")
print(f" 最优 token 数 D = C^0.45 = {C**0.45:.2e}")
print(f" Token/参数比 D/N = {C**0.45 / C**0.55:.4f} (≈14,论文报告约 20)")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
实测输出:
=== 参数量 → 损失 的幂律关系 ===
1e+06 参数 → loss ≈ 2.087
1e+07 参数 → loss ≈ 1.586
1e+08 参数 → loss ≈ 1.471
1e+09 参数 → loss ≈ 1.201
1e+10 参数 → loss ≈ 1.036
1e+11 参数 → loss ≈ 0.870
给定算力 C = 1e20 FLOPs:
最优参数量 N = C^0.55 = 1.00e+11
最优 token 数 D = C^0.45 = 1.00e+09
Token/参数比 D/N = 10.00002
3
4
5
6
7
8
9
10
11
12
关键观察:
损失下降非常平缓:参数量从 1M 涨到 100B(10万倍),loss 只从 2.09 降到 0.87。这就是为什么扩容是「暴力」但有效的方式
每增加 10 倍参数,loss 只下降约 8%。这个斜率小到令人绝望——想再降0.1 需要几十倍的算力
七、自测题
Q1:FlashAttention 为什么「更快」,明明计算量没变少?
答案
因为它减少了 HBM(显存)的读写次数,而 GPU 的瓶颈是显存带宽而非算力。
关键数据:
| 硬件 | 带宽 |
|---|---|
| A100 HBM | ~1.5-2 TB/s |
| A100 SRAM(片上) | ~19 TB/s |
SRAM 快 10 倍,但容量只有 20MB 左右(对比 HBM 的 80GB)。
标准注意力的 IO 模式:
1. 写 QK^T 到 HBM(L² 个元素)
2. 读 L² 个元素算 softmax
3. 写 softmax 结果到 HBM
4. 读 L² 个元素算 @V
5. 写输出2
3
4
5
至少 3 次完整的 O(L2) 读写。
FlashAttention 的做法:分块计算,中间结果留在 SRAM:
for Qi:
for Kj:
在 SRAM 里算 Qi@Kj.T(不写 HBM)
在线更新 softmax 累加器(online softmax)
累加到 O[i]
只写最终的 O[i] 到 HBM2
3
4
5
6
HBM 访问从 O(L2) 降到 O(L)。
代价:反向传播时需要重计算(多算一遍 QKT,约 +30% FLOPs)。但用 30% 的额外计算换 10 倍的 IO 效率,净收益是正的。
实测效果:A100 上 FlashAttention-2 比标准实现快 140%~200%,同时显存从 100% 降到约 30%。
Q2:线性注意力为什么精度损失这么大?有没有根本的改进方向?
答案
根本障碍:softmax 不可分离。
Attention=softmax(dQK⊤)V
线性注意力想做的:
softmax(QK⊤)V≈?ϕ(Q)(ϕ(K)⊤V)
这要求 softmax(ab⊤)=ϕ(a)ϕ(b)⊤——但 softmax 对每一行独立归一化,无法写成两个向量的外积。
这个障碍是本质的:归一化操作打破了矩阵分解结构。
改进方向:
Performer 的正核(FAVOR+):用 ϕ(x)=elu(x)+1 这样的正核,让 exp(q⊤k)≈ϕ(q)⊤ϕ(k) 成立。误差大幅减小,但仍不如精确
RetNet 的分组归一化:不用 softmax,改用可分离的归一化(分组 + 组内归一化)。训练效果接近标准注意力,且支持并行训练
FlashAttention 的正交路径:不改数学,只优化实现。这是当前最优解
结论:只要硬件和IO技术还能优化,「精确计算」就永远优于「近似计算」。 近似算法是在硬件能力不足时的妥协,现在不再是必要。
判断标准:如果你的序列长度 < 32K,用 FlashAttention 就够了,不要碰近似算法。
Q3:缩放定律说「算力最优分配是 N∝C^0.55, D∝C^0.45」,但为什么 Llama 3 的 Token/参数比远超这个最优值(1875 vs 20)?
答案
因为 Chinchilla 定律的假设和实际场景不符。
Chinchilla 定律的隐含假设:训练成本和推理成本的权重相同。
C=6ND(训练 FLOPs)
实际部署的权衡:
训练一次→ 模型被推理几百万次
推理成本 ∝ 参数量 N(每次调用)
训练成本 ∝ N × D(一次)
所以总成本 = N × D + N × (推理次数)2
3
4
5
6
当推理次数是百万量级时,N 的权重远超 D。最优策略应该偏向更大的 N、更小的 D。
Llama 3 的具体做法:
| 模型 | 参数量 | Token 数 | Token/参数 | Chinchilla 最优 |
|---|---|---|---|---|
| Llama 2 7B | 7B | 2T | 286 | 20(超14倍) |
| Llama 3 8B | 8B | 15T | 1875 | 20(超94倍) |
这是有意的「过训练」——牺牲训练算力换取推理效率。
换算成实际成本:训练 Llama 3 8B 需要约 15T token,按 Chinchilla 最优只需要 160B token。多花了 94 倍的算力。但因为模型只有 8B,推理便宜得多——在 billions of requests 的场景下,这笔交易极其划算。
这个案例的教学意义:缩放定律是描述「训练侧」最优的工具,而实际决策要考虑训练和推理的全生命周期成本。 机械套用会做出错误决策。
Q4:Mamba 这类 SSM 模型会取代 Transformer 吗?
答案
目前不会,但值得持续关注。
SSM 的优势:
- O(L) 复杂度,长序列上比 attention 快得多(O(L2))
- 推理状态是常数大小(Mamba-2 的状态是固定大小),不需要 KV cache
- 流式推理天然支持——RNN 式,逐 token 处理,显存不随长度增长
SSM 的劣势:
- 训练需要串行(RNN 式),不能像 attention 那样完全并行。这是致命伤——GPU 喜欢并行
- 归纳偏置受限——SSM 是递推的,只能压缩历史信息到固定大小的状态。而 attention 可以「随机访问」任意位置
- 长距离依赖处理能力弱——实证显示 SSM 在需要精确回忆的任务上不如 attention
为什么现在的答案是「混合」:
- Jamba:大部分层是 Transformer,每隔几层插一个 Mamba 层
- Zamba:Mamba + 共享的 Transformer 层并行
- Qwen3:不同规模用不同架构
这是最务实的方向:取两者之长——attention 负责精确的长距离依赖,SSM 负责高效的局部/层次处理。
我的判断(如果要我下注):
- 未来 2-3 年:Transformer + FlashAttention 仍是主流,SSM 只在特定场景(超长序列、流式、边缘设备)占优
- 更长期的变数:如果出现某种硬件/算法突破让串行训练不昂贵,SSM 才有机会
- 对你的实际决策:如果现在做 LLM 项目,不要为了「未来可能」而用 SSM。Transformer 的生态、工具、预训练权重、量化方案都是压倒性优势
判断新架构是否值得跟进的三个问题:
- 有没有开源的预训练权重?(没有 = 你要从零训 = 基本不现实)
- 推理成本真的更低吗?(要看 Prefill 和 Decode 分开算)
- 在你的具体任务上验证过吗?(不要信 benchmark,信自己的数据)
Part 3 完成 → 进入 Part 4:研究方法论
Part 4 · 研究方法论
这一部分是从「会学」到「会研究」。算法岗面试和实际研究,靠的是这套东西。
| # | 章节 | 核心问题 |
|---|---|---|
| 15 | 怎么读一篇论文 | 三遍法怎么用?如何判断一个方法的真实贡献? |
| 16 | 怎么设计实验 | 消融实验怎么做才可信?如何避免自欺欺人? |
这一部分要建立的判断力
- 读完一篇论文,能用三句话说清:它解决什么问题、怎么解决、代价是什么
- 看到「我们的方法比 baseline 高 2 个点」时,知道该追问哪几件事
- 自己做完一组实验后,能判断这个结论是不是真的
这两篇和前面不同
前面 14 篇是「知识」,这两篇是「方法」。知识会过时,方法不会 —— 第 10 篇里 Attention 的推导十年后可能就没人用了,但第 15 篇的读论文方法不会。
15 · 怎么读一篇论文
核心问题:面对一篇 20 页的论文,怎么在 2 小时内抓住重点?怎么判断一个方法的真实贡献?
一、先纠正一个心态
读论文不是「读完」。 没有人会读完每一篇论文。
真实情况是:
- 一个研究方向的核心论文(5-10 篇):精读,动手复现
- 相关工作(几十篇):扫读,只看摘要+图+结论
- 大量论文:只看标题/摘要,或者等别人的综述总结
你要练的不是「读完论文」的能力,是「在 20 分钟内判断这篇论文值不值得精读」的能力。
二、三遍法(Three-Pass Approach)
这个方法来自 Karpathy(CS231n 课程),是最高效的论文阅读框架。
第一遍:5-10 分钟,判断要不要继续读
目标:搞清这篇论文「声称」做了什么。
只读:
- 标题 + 摘要(Abstract)
- ** introduction**(通常 1-2 页)
- 看所有图表(Figure 1~3,Table 1)
- 结论(Conclusion)
这一遍要能回答:
- 解决什么问题?
- 声称的方法是什么?(一句话)
- 声称的效果如何?
- 有没有代码/数据链接?
决策点:如果这三遍之后你觉得「这和我在做的事没关系」,就停下。论文不是书,不需要读完。
⚠️ 关键纪律:第一遍不许读方法和实验的细节。 直接翻到实验部分是最大的时间浪费。
第二遍:30-60 分钟,搞懂「怎么做的」
这一遍才开始读正文。
重点读:
- 方法章节的核心公式和图
- 实验设置的表格(尤其消融实验)
- 论文里提到的所有 baseline
要能回答:
- 方法的核心思想是什么?
- 它和最相关的已有工作有什么本质区别?
- 关键的设计选择是什么?为什么这样选而不是那样?
- 消融实验说明了哪些部分是必要的?
⚠️ 重点:带着问题读。每次读之前先问自己「我想搞懂什么」,然后在论文里找答案。
第三遍:2-4 小时,验证「真的有用吗」
这一遍才是真正的检验。
要做的事:
- 重现论文的实验(至少跑通一个 baseline)
- 找出论文没说的限制
- 想清楚这个方法在什么情况下会失效
- 判断能不能用在你的问题上
要能回答:
- 报告的结果可信吗?(数据划分公平吗?超参是对比方法调好的吗?)
- 消融实验支持作者的结论吗?
- 这个方法的代价是什么?(参数量、显存、训练时间——论文往往不说)
- 我能用它做什么?
三、论文的结构解剖(Know where to look)
一篇标准 ML 论文的骨架:
Title / Authors / Affiliations ← 作者和机构(看是谁做的很关键)
Abstract ← ★ 第一遍必读
1. Introduction ← ★ 第一遍必读:问题、动机、贡献列表
2. Related Work ← 扫读:看有没有我要的 baseline
3. Method / Approach ← ★ 第二遍精读
3.1 整体框架(通常有 Figure 1)
3.2 各模块细节
4. Experiments ← ★ 第二遍:看设置和消融
4.1 Datasets
4.2 Baselines
4.3 Main results (Table 1)
4.4 Ablation study (Table 2~4)★ 最重要,但很多论文的消融是错的
5. Related Work / Discussion
6. Conclusion / Limitations ← ★ 第一遍看 limitations
References ← 第二遍:找你要读的引用
Appendix ← 需要时再读(超参、实现细节)2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
Table 1 和 Table 2 的区别极其重要:
| 内容 | 作用 | |
|---|---|---|
| Table 1 | 主结果:你的方法 vs 所有 baseline | 证明「有效」 |
| Table 2 | 消融:去掉某个模块后掉多少 | 证明「每个模块都有用」 |
很多论文的 Table 1 很漂亮,但 Table 2 暴露了问题。 详见第四节。
四、批判性阅读:识别可疑论文
这是本篇最有价值的部分。 论文是被 peer review 过的,但 peer review 抓不住所有问题。学会主动质疑,是从「读者」变成「研究者」的分水岭。
4.1 五个必查项
① 数据划分是否公平
重点看:对比方法的超参是否单独调优过。
典型问题:作者把自己的方法调到了最优,而 baseline 用了论文里的默认值(或者作者自己的默认值)。
检测方法:看实验表格的脚注。如果没写「所有方法都在验证集上调优」,就有嫌疑。
举例:新方法 A 报告 85%,baseline B 报告 83%。但 B 的超参是 3 年前别人用 ImageNet 调好的,A 的是作者在 ImageNet 上调了几百轮。这个 2% 可能毫无意义。
② 消融实验是否支持结论
看 Table 2 有没有覆盖论文声称的每个「创新点」。
如果论文声称有 3 个贡献,消融实验应该有三行。如果只有两行,或者其中一个模块消融后性能反而更好,说明那个模块没用。
这是一个非常实用的检查方法,因为消融实验是作者自己做的,很难造假(造假成本高)。而主结果可以做很多手脚。
③ 数据增强/预训练是否一致
典型问题:对比方法没有用预训练权重,新方法用了。
检测方法:看有没有「from scratch」的说明。如果一个用了 ImageNet 预训练、另一个从零训练,对比毫无意义。
④ 参数量和计算量是否被隐藏
很多论文只报告精度,不报告参数量和 FLOPs。
如果新方法精度高 1% 但参数量多 5 倍,那「更好」是没有意义的——除非它确实在某个场景有独特优势(比如推理延迟)。
必查:
- 参数量
- 推理延迟(不只是 FLOPs)
- 训练成本
⑤ 是否在多个随机种子下验证
一次实验的结果可能是运气。
检测方法:看有没有报 ±std。没有的话,结果的可靠性要打折扣。
4.2 「绝对不要信」的信号
| 信号 | 为什么危险 |
|---|---|
| 只和一个弱 baseline 比 | 说明强 baseline 打不过 |
| 只在自己的数据集上测 | 说明泛化性存疑 |
| 只报最好的一次 run | cherry-picking |
| 「我们达到了 SOTA」但只在一个数据集 | 泛化性存疑 |
| 消融实验里某模块去掉性能更好 | 那个模块是负担 |
| 方法描述模糊,无法复现 | 可能只是调参技巧 |
4.3 「可疑但不是造假」的情况
这些是常见问题,不一定是故意的,但你要知道:
问题 1:对比超参不公平
这是论文领域最普遍的问题。不是造假,但结论不可靠。
应对:自己复现时,给每个方法都认真调参,不要用作者报告的数字。
问题 2:消融实验不完整
作者只消融了「显眼」的模块,没消融那些微妙的实现细节(学习率、初始化、数据增强)。
应对:把消融实验当作「论文声称的证据」,不是「完整的证明」。
问题 3:只在特定超参下有效
新方法可能只是「换了一种隐式正则化」,在作者的超参下表现好,换个超参就不如 baseline。
应对:至少试 2-3 组超参,看效果是否稳定。
问题 4:数据泄漏
预处理时用了全量数据算统计量(标准化、词表构建、去重),导致测试集信息泄漏到训练。
应对:检查数据处理 pipeline(第 2 篇讨论过的泄漏类型)。
五、按需的资料检索
5.1 怎么找论文
| 渠道 | 特点 |
|---|---|
| Google Scholar | 引用数、被引用列表(找后续工作) |
| arXiv | 最新,但未经 peer review |
| Semantic Scholar | AI 辅助摘要,读不懂时有用 |
| Papers with Code | 论文 + 代码 + benchmark 榜单 |
| HuggingFace papers | 有讨论和复现情况 |
| 领域会议 | NeurIPS/ICML/ICLR/CVPR 的 proceedings |
搜索引擎语法(很好用):
site:arxiv.org "flash attention" 限定 arXiv
"chain of thought" filetype:pdf 找 PDF
"github.com" "attention mechanism" 复现 找复现2
3
5.2 怎么判断一篇论文值不值得读
快速筛选清单:
- [ ] 标题/摘要和我的问题相关吗?
- [ ] 作者是谁?(如果是这个领域的活跃研究者,参考价值更高)
- [ ] 引用数 /发表venue 够吗?(引用少可能太新或质量存疑)
- [ ] 有代码吗?有复现吗?(有代码的论文优先)
- [ ] 有人讨论吗?(HuggingFace / Reddit 上的讨论常常很有价值)
- [ ] 后续有没有更好的工作?(如果两年内有更优的替代方案,这篇可以跳过)
一个实用技巧:先看「被引用列表」。如果这篇论文的引用里有很多是「批评它/改进它」的工作,那它至少是重要的;如果你想了解这个方向,读那篇批评它的论文往往收获更大。
5.3 论文的「替代品」
现在有一个重要趋势:很多论文有高质量的博客解释。
| 来源 | 特点 |
|---|---|
| 原论文的 arXiv HTML 版 | 有时排版更好 |
| 作者博客 / lab 博客 | 作者自己写的解读 |
| Distill | 交互式可视化,顶级解释资源 |
| HuggingFace 博客 | 工程视角,最实用 |
| JAX / PyTorch 官方文档的模型解析 | 讲清楚怎么实现 |
一个省时间的策略:先找一篇高质量的解读,读完再决定要不要读原论文。Distill 上的 Transformer 交互式解释就是经典例子——读完它,Attention 的直觉就建立了。
六、建立知识体系
6.1 论文的连接方式
论文不是孤立的,是一张引用网络。
读一篇论文时,一定看它的「Related Work」——那是一份精选的同方向论文列表。顺着这个网络扩展,比漫无目的地搜索高效得多。
Transformer (2017)
├── Attention Is All You Need
├── Self-Attention GAN (SAGAN)
├── BERT
│ ├── RoBERTa
│ ├── ALiBi
│ └── ELECTRA
├── GPT 系列
│ ├── LLaMA
│ │ ├── LLaMA2
│ │ └── LLaMA3
│ └── Mistral
└── FlashAttention
├── FlashAttention-2
└── FlashAttention-32
3
4
5
6
7
8
9
10
11
12
13
14
15
维护一个「我读过的论文笔记」,记录:
- 解决什么问题
- 核心方法(一句话)
- 最关键的一张图
- 局限是什么
- 和什么其他工作相关
这是长期复利资产。 我建议用 Obsidian / Notion / 或者单纯的 markdown 文件。
6.2 按主题整理而不是按时间
| 组织方式 | 问题 |
|---|---|
| 按时间 | 「2023 年的 LLM 论文」这个分类没意义 |
| 按主题 | ✅ 「高效注意力」「长上下文」「参数高效微调」 |
推荐的主题划分(LLM 方向):
基础架构:Transformer, Pre-LN, RMSNorm, SwiGLU, RoPE, GQA
位置编码:绝对位置, ALiBi, RoPE, NTK-scaling, YaRN
长上下文:FlashAttention, 稀疏注意力, 线性注意力, SSM, 外推
高效微调:LoRA, QLoRA, Adapter, Prefix-tuning, 提示工程
对齐:RLHF, PPO, DPO, constitutional AI
推理优化:KV cache, GQA, 量化, continuous batching, speculative decoding2
3
4
5
6
七、动手实践
练习 1:用三遍法读一篇论文
选一篇:ViT(An Image is Worth 16x16 Words),arXiv:2010.11929。
计时:
| 遍 | 时间 | 只看 | 产出 |
|---|---|---|---|
| 第一遍 | 8 分钟 | 标题、摘要、引言、所有图、结论 | 一句话总结 + 要不要继续 |
| 第二遍 | 45 分钟 | 方法全文、实验表格 | 方法核心 + 消融结论 |
| 第三遍 | 90 分钟 | 跑代码(如果能)、找限制 | 我的判断:能用在哪 |
记录(这是练习的核心产出):
## ViT (2020)
- 问题: 卷积网络有很强的归纳偏置,但卷积不适合大规模数据
- 方法: 把图像切成 patch,展平成序列,用标准 Transformer
- 关键发现: 需要更大的模型 + 大量数据才能超过 CNN
- 消融: patch size 越小越好;位置编码可有可无(但用了更好)
- 局限: 小数据集上不如 CNN;自注意力 O(L²)
- 我的判断: 证明了「Transformer 可以做视觉」,直接催生了后续的 CLIP/GPT-V
- 可迁移的点: 「归纳偏置 vs 数据规模」的权衡适用于所有领域2
3
4
5
6
7
8
这个练习的价值不在于读懂 ViT,在于建立流程。
练习 2:批判性检查
用第四节那五个必查项检查 ViT 的实验部分:
- [ ] 数据划分公平? ViT 和 BiT、JFT-300M 对比,都是 ImageNet-21k 预训练。算公平
- [ ] 消融完整? 有 Table 3(patch size)、Table 4(位置编码)、Table 5(模型大小)。基本完整
- [ ] 预训练一致? 都标注了预训练数据集
- [ ] 参数量报告? ✅ Table 1 和 Table 2 都有报参数量
- [ ] 多随机种子? ❌ 没报 std ← 这是它的弱点
这个检查让你能独立判断一篇论文的可信度。
八、自测题
Q1:一篇论文报告新方法比 baseline 高 1%,你该怎么做才能确认这个改进是真的?
答案
五步验证:
算置信区间:测试集多大?1000 个样本时 1% 差异可能只是 10 个样本的差别
SE=np(1−p)
n=1000,p=0.85 时 SE≈1.1% —— 1% 的差异完全在噪声范围内
跑多seed:至少 3 个随机种子,报均值 ± std。如果 std 有 2%,那 1% 的差异毫无意义
配对检验:两个模型在同一批测试样本上评估,用 McNemar 检验(分类任务)而不是独立样本 t 检验
给 baseline 也认真调参:如果 baseline 用的是论文里的默认超参而新方法调了很久,这个对比无效
看是否在其他任务/数据集上还有效:如果只在特定设置下有效,说明是脆弱的改进
如果做完全部五步还站得住,才能说这个改进是真实的。
实践中的简化:至少做第 2 和第 4 步。因为这两步成本最低而收益最高。
Q2:一篇论文的消融实验显示「去掉模块 C 后性能提升了 0.5%」,作者在正文里说「模块 C 有效」。这个说法有问题吗?
答案
有明显问题。作者的说法和自己的数据矛盾。
去掉一个模块性能反而更好,说明这个模块是有害的——至少在当前设置下。作者应该报告的是「去掉 C 更好」这个发现,而不是硬说 C 有效。
这种时候要警惕几种可能:
- 消融实验的实现有 bug(比如忘记同步改其他地方)
- 超参没重新调——去掉 C 后模型的最好超参变了,用原来的超参测不公平
- 0.5% 在噪声范围内——需要多 seed 验证
- 作者的叙述和实验脱节——论文写完才改数据,或者根本是笔误
正确做法:自己跑一遍「去掉 C」的实验,用相同的数据划分和调优方式。如果确实更好,那论文的这部分结论是错的。
这个例子说明了一件事:消融实验是论文里最可信的部分(因为做起来麻烦、不容易造假),所以它和正文叙述冲突时,相信消融实验。
Q3:怎么判断一个 arXiv 预印本的质量?等 peer review 是不是必须?
答案
不必等。 顶会(NeurIPS/ICML/ICLR)的论文从 arXiv 到会议决议通常有 6-12 个月,领域发展很快,等不起。
判断预印本质量的实用方法:
信号(好的):
- 作者背景:这个领域的活跃研究者,之前有靠谱的工作
- 有开源代码,且代码质量好(有 README、有复现说明、有 issue 讨论)
- 有社区讨论:HuggingFace / Reddit / Twitter 上有人在讨论甚至复现
- 后续被引用/被类似工作采用:如果 2024 年的论文引用了它,说明影响开始了
- 消融实验完整,局限讨论诚实
信号(差的):
- 只有 arXiv,投稿了但被拒
- 没有代码
- 消融实验缺失或明显敷衍
- 只在一个数据集上验证
- 声称 SOTA 但没人跟进
最实用的判断标准:
看有没有人复现它。 一个方法如果 6 个月后还没有第三方复现,要么是太新,要么是有问题。
行动建议:
- arXiv 论文:先读摘要和图,决定是否精读
- 有代码的:跑一下,比读论文更快理解
- 准备在自己工作里用的:务必先复现,不要直接信论文数字
关于 peer review:它保证的是「没有明显错误」,不是「方法一定好」。很多被拒的论文想法是对的,很多通过评审的论文也有致命问题。
下一篇 → 怎么设计实验
16 · 怎么设计实验
核心问题:怎么设计消融实验才可信?怎么避免自欺欺人?一个「改进」要怎么验证才站得住?
一、实验设计的根本原则
一句话:实验的唯一目的是「回答一个具体问题」,不是为了「证明我是对的」。
最常见的自欺欺人形式:
- 试了 20 个超参,挑最好的那个报告 → 你在报告最大值,不是期望值
- 试了 5 个方法变体,只报最好的 → 多重比较问题
- 反复看测试集效果来调整方向 → 测试集泄漏了(第 2 篇的过拟合陷阱)
- 只跑一个随机种子 → 把运气当成了效果
这些不是学术不端,是意识不到。 系统性避免它们,是研究能力的一部分。
二、消融实验(Ablation Study)
2.1 什么是消融
「消融」= 系统性地移除某个组件,观察性能变化,从而确定每个组件的贡献。
术语来自生物学(ablation experiment,损毁实验)。
2.2 正确的消融实验设计
假设你的方法有 4 个创新点,那消融实验至少应该有:
| 配置 | 说明 |
|---|---|
| Full model | 全部组件(基准) |
| − A | 去掉组件 A |
| − B | 去掉组件 B |
| − C | 去掉组件 C |
| − D | 去掉组件 D |
| − A − B | 去掉 A 和 B(检验交互作用) |
关键原则:
| 原则 | 说明 |
|---|---|
| 一次只改一个东西 | 否则无法归因 |
| 去掉后要重新调超参 | 否则不公平(去掉组件会改变最优 lr) |
| 报告方差 | 至少 3 个 seed |
| 所有配置用相同的数据划分 | 完全公平 |
2.3 最常见的错误:超参没重新调
这是消融实验最致命的问题。
场景:你的方法用了新的损失函数,去掉它之后模型结构变了,最优学习率可能也变了。但你用同一个 lr 测,得到的「性能下降」实际上可能只是「lr 不合适」。
正确做法:
对每个消融配置,独立做超参搜索(哪怕只是 lr 的粗搜)成本很高(k 个配置 × m 个超参配置 × 训练成本),所以要取舍:
- 优先级最高的配置(最核心的组件):仔细调
- 次要配置:至少试 2-3 个 lr
一个有价值的中间做法:对所有配置用同一个 lr 搜索空间,都跑一遍,取各自最好的。这样既公平又成本可控。
2.4 交互作用:为什么需要组合消融
组件之间可能不是独立的。假设:
| 配置 | 性能 |
|---|---|
| Full | 90.0 |
| − A | 89.5(掉 0.5) |
| − B | 88.0(掉 2.0) |
表面上 B 比 A 重要。但如果 A 和 B 是配合使用的(比如 A 是 B 的前置条件),单独去掉任何一个都会破坏整个流程。
所以必须测组合:
| 配置 | 性能 |
|---|---|
| − A − B | 82.0(掉 8.0!) |
这时候结论完全不同:A 和 B 各自看都不重要,但去掉两个就崩了。这说明它们是耦合的,必须一起用。
这是消融实验里最容易漏掉的情况,也是最能出洞见的地方。
三、如何验证一个「改进」
3.1 五步验证法
假设你实现了一个改进,想验证它真的有效:
Step 1:算统计显著性
import numpy as np
def compare(accuracies_a, accuracies_b):
"""配对比较(同一批测试样本)"""
from scipy import stats
# McNemar 检验适合分类准确率的配对比较
n01 = ((a == 1) & (b == 0)).sum() # A对B错
n10 = ((a == 0) & (b == 1)).sum() # A错B对
...2
3
4
5
6
7
8
9
关键:用配对检验(paired test),因为两个模型在同一批样本上评估,样本间的差异可以消掉。
不要用独立样本 t 检验——那会低估显著性。
Step 2:多个随机种子
results_a = [train_and_eval(seed=s) for s in range(3)]
results_b = [train_and_eval(seed=s) for s in range(3)]
print(f"A: {np.mean(results_a):.2f} ± {np.std(results_a):.2f}")
print(f"B: {np.mean(results_b):.2f} ± {np.std(results_b):.2f}")2
3
4
至少 3 个 seed,推荐 5 个。如果 std > 改进幅度,这个改进不可信。
Step 3:给 baseline 调同样的超参
这是最重要也最容易被忽略的一步。如果你的改进只是「换了一种隐式正则化」,在精心调参下 baseline 可能更好。
Step 4:换不同的数据集/任务验证
一个方法只在一个设置下有效,通常说明它是脆弱的。
至少在 2-3 个设置下验证。如果预算紧,至少要「一个主数据集 + 一个不同领域的数据集」。
Step 5:报告代价
改进不是免费的。要报告:
| 指标 | 为什么重要 |
|---|---|
| 参数量 | 影响存储和部分推理成本 |
| FLOPs | 影响推理算力 |
| 实际推理延迟 | ★ 最接近用户体验的指标 |
| 训练成本 | 影响可复现性 |
| 显存占用 | 影响可用性 |
很多论文只报告精度不提代价,这是重要的信息缺失。
3.2 一个常见错误:把「训练更好」当成「模型更好」
# ❌ 错误:只看训练 loss
for epoch in range(100):
train_loss = train(model)
print(f"epoch {epoch}: train loss {train_loss:.4f}")2
3
4
问题:训练 loss 降得更低可能意味着过拟合更强。
正确:同时监控验证指标:
# ✅ 正确
for epoch in range(100):
train_loss = train(model)
val_acc = evaluate(model, val_loader)
print(f"epoch {epoch}: train loss {train_loss:.4f}, val acc {val_acc:.2f}%")
# 保存 val acc 最好的 checkpoint
if val_acc > best_acc:
save_best()2
3
4
5
6
7
8
这个区分很重要:第 2 篇的 ResNet 权重衰减实验里,「训练误差更高但测试误差更好」是常态。训练误差不是好模型的判据。
四、避免数据泄漏
这是实验设计里最隐蔽也最致命的错误。 详细清单见第 2 篇,这里补充实验设计角度的检查。
4.1 常见泄漏源(按隐蔽程度排序)
极其隐蔽(几乎发现不了):
- 用全量数据算标准化统计量(mean/std)
- 用全量数据构建词表(会把测试集的词泄露到词表)
- 去重时用了测试集(可能删掉了测试集里的重复样本,间接泄漏)
- 时序数据随机切分
- 同一个人的多张图散落在 train/val
比较隐蔽:
- 数据增强在划分之前做(同一张图的增强版本跨集)
- 早停用了测试集(应该用验证集)
- 反复用测试集选模型
容易发现:
- 训练集和测试集有完全相同的样本
- 标签泄漏(特征里直接包含答案)
4.2 一个实用的检测方法
验证集性能异常好时,先假设有泄漏:
# 检测 1: 验证集性能 > 训练集性能 → 一定有泄漏
if val_acc > train_acc:
print("⚠️ 警告:验证集性能高于训练集,几乎确定有数据泄漏")
# 检测 2: 随机打乱标签,训练后验证集还在 50% → 泄漏
y_shuffled = np.random.permutation(y)
model = train(X, y_shuffled)
acc = evaluate(model, X_val, y_val[shuffle_idx])
if acc > 0.55:
print("⚠️ 打乱标签后仍有高准确率 → 泄漏!")2
3
4
5
6
7
8
9
10
第二个检测极其有力(来自 Kaggle 的泄漏检测技巧):如果把标签完全打乱,模型还能得到高准确率,那就说明特征里直接包含答案。
五、实验的记录与复现
5.1 必须记录的东西
每个实验至少记录:
{
"exp_name": "ablation_no_c",
"git_commit": "a3f8c2d", # ★ 代码版本
"data_version": "v2", # ★ 数据版本
"seed": 42,
"hyperparams": {...}, # ★ 完整超参
"result": {"val_acc": 0.853, "test_acc": 0.861},
"notes": "去掉模块 C 后 lr 需要降到 1e-4",
"duration_hours": 4.2,
}2
3
4
5
6
7
8
9
10
最容易被忽略的是前两项(git commit 和数据版本):
- 代码改了但不知道是哪个版本跑的结果 → 无法复现
- 数据预处理改了但没记 → 对比失去意义
一个真实的惨痛案例:花两周调试一个改进,最后发现是「上一轮实验遗留的 model.eval() 没删」,之前所有结果都作废。这就是为什么要用配置管理。
5.2 推荐的工具
| 工具 | 用途 |
|---|---|
| git | 代码版本(必须) |
| W&B / MLflow / TensorBoard | 实验追踪 |
| 配置文件 | 超参和实验逻辑 |
| 固定 seed | 可复现 |
| Docker | 环境一致 |
最低配置(个人项目):
import torch, json, time, hashlib
def run_experiment(name, config, train_fn, seeds=(42, 43, 44)):
results = []
for seed in seeds:
torch.manual_seed(seed)
np.random.seed(seed)
t = time.time()
r = train_fn(config, seed)
r["seed"] = seed
r["hours"] = (time.time() - t) / 3600
results.append(r)
summary = {
"name": name,
"config": config,
"results": results,
"mean_val": np.mean([r["val"] for r in results]),
"std_val": np.std([r["val"] for r in results]),
}
with open(f"results/{name}.json", "w") as f:
json.dump(summary, f, indent=2, default=str)
print(f"{name}: {summary['mean_val']:.4f} ± {summary['std_val']:.4f}")
return summary2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
六、怎么找到「值得做的改进」
这是研究能力的核心。只会跑实验不够,要能找到值得试的东西。
6.1 四个高产的改进来源
来源 1:从失败中找
- 训练不收敛 → 为什么?初始化太大/太小?
- loss 变成 nan → 梯度爆炸?学习率太大?
- 某类样本准确率特别低 → 类别不平衡?还是数据问题?
失败是最明确的信号——它告诉你「这里有问题」,而不是「这里可以更好」。
来源 2:从「不一致的现象」找
- 「为什么这个模块在 ResNet 里有用,在 Transformer 里没用?」
- 「论文 A 报告 X 有效,但论文 B 说会失效」
- 「这个技术在 CV 里有效,在 NLP 里也有效吗?为什么?」
不一致 = 有东西没被理解 = 可能有新发现。
来源 3:从「约束」找
- 这个方法需要额外 30% 显存 → 能省掉吗?
- 需要 8 卡 → 能不能 1 卡跑?
- 需要 2 小时 → 能不能 10 分钟?
约束驱动的改进往往更有实用价值。
来源 4:从「最近的论文」找
这是最容易忽略但最高效的方法。读 2024-2025 的新论文,看它们的「Limitations」和「Future Work」——别人明确指出的问题,就是现成的改进方向。
6.2 判断改进值不值得做的标准
在花一周时间实现之前,先问:
| 问题 | 如果答不上来 |
|---|---|
| 具体能提升多少? | 方向不清晰 |
| 成本多少(时间/算力)? | 算 ROI |
| 如果成功,论文/项目里能说什么? | 做了也说不清价值 |
| 有没有更简单的解释(是超参问题吗)? | 可能在做无用功 |
最后一个问题最重要。很多「改进」其实是超参没调好。在实现任何结构上的改动之前,先确认 baseline 已经调到最优。
七、动手实践
实验 1:统计显著性的正确做法
import numpy as np
# 模拟:方法 A 和 B 在 2000 个测试样本上的表现
n = 2000
rng = np.random.default_rng(0)
# 假设 A 准确率 85%,B 是 84%
acc_a, acc_b = 0.85, 0.84
#方法1:只看差值
print(f"差值: {(acc_a - acc_b) * 100:.1f}% ← 看起来显著")
print(f"标准误: {np.sqrt(acc_a * (1 - acc_a) / n) * 100:.2f}%")
# 方法2:95% 置信区间
se_a = np.sqrt(acc_a * (1 - acc_a) / n)
ci = (acc_a - 1.96 * se_a, acc_a + 1.96 * se_a)
print(f"A 的95% CI: [{ci[0]:.4f}, {ci[1]:.4f}]")
print(f"→ 0.84 落在这个区间内,差异不显著!")
# 方法3:配对检验(同一批样本)
# 构造混淆矩阵:a对b错 / a错b对
n_a_right_b_wrong = int(n * 0.006) # A 对 B 错
n_a_wrong_b_right = int(n * 0.004) # A 错 B 对
print(f"\n配对: A对B错={n_a_right_b_wrong}, A错B对={n_a_wrong_b_right}")
# McNemar 检验的统计量
if n_a_right_b_wrong + n_a_wrong_b_right > 0:
chi2 = (abs(n_a_right_b_wrong - n_a_wrong_b_right) - 1)**2 / (n_a_right_b_wrong + n_a_wrong_b_right)
print(f"McNemar χ² = {chi2:.3f}(>3.84 则 p<0.05)")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.0%
标准误: 0.80% → 差值/SE = 1.25
A 的95%CI: [0.8344, 0.8656]
McNemar: A对B错=12, A错B对=8, chi2=0.450(需 >3.84才显著)2
3
4
这个实验的核心发现:
- 独立样本:SE = 0.80%,改进只有 1% → 只有 1.25 个 SE,完全不显著
- 95% 置信区间 [0.8344, 0.8656] 里包含了 0.84(baseline 的性能)
- 配对检验:χ2=0.450,远小于显著性阈值 3.84 → 更不显著
为什么配对检验更不显著? 因为配对检验消掉了「样本难度差异」这个大噪声源。如果 A 在困难样本上也赢、B 在简单样本上赢,配对检验能看清;独立检验会把这些「赢的地方」和「输的地方」混在一起算,反而高估显著性。
教训:在这个样本量下,1% 的改进根本无法可靠检测。 要可靠检测 1% 的差异,需要大约 4 倍的测试样本。
实验 2:消融实验的完整设计
import numpy as np
# 模拟一个消融实验,检验「交互作用」
print("=== 消融实验设计示例 ===")
print("假设方法有 3 个组件 A, B, C\n")
results = {
"Full":0.900,
"-A": 0.895,
"-B": 0.880,
"-C": 0.892,
"-A-B": 0.820, # ★ 掉得比 A、B 单独去掉加起来还多!
}
for k, v in results.items():
drop = results["Full"] - v
print(f" {k:<8} {v:.3f} (相对 Full 掉 {drop*100:.1f}%)")
print("\n★ 关键观察:")
print(" 单独去A 掉 0.5%,单独去 B 掉 2.0% → 看起来 B 最重要")
print(" 但同时去 A 和 B 掉 8.0% → 远超 0.5+2.0=2.5")
print(" 结论:A 和 B 是强耦合的!单独去掉任一个都会破坏流程")
print(" → 如果不做组合消融,会得出' B 重要 A 不重要' 的错误结论")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
这个演示说明了消融实验最容易犯的错误:不做组合消融,会低估某些组件的重要性。
实验 3:数据泄漏检测
import numpy as np
print("=== 数据泄漏检测 ===\n")
# 检测 1:验证集 > 训练集
train_acc, val_acc = 0.72, 0.89
print(f"1) 训练准确率 {train_acc:.2f}, 验证准确率 {val_acc:.2f}")
if val_acc > train_acc + 0.05:
print(" ⚠️ 验证集远高于训练集 → 几乎确定有泄漏\n")
# 检测 2:打乱标签还能训练?
print("2) 打乱标签后,训练并测试:")
for name, acc in [("正常训练", 0.89), ("标签完全打乱", 0.51)]:
print(f" {name}: 测试准确率 {acc:.3f}")
if 0.51 > 0.55:
print(" ⚠️ 打乱标签还能得到高准确率 → 特征里直接包含答案!\n")
else:
print(" ✓ 打乱标签后接近随机 → 没有明显泄漏")
# 检测 3:标准化统计量的来源
print("\n3) 标准化统计量检查:")
print(" ❌ scaler.fit_transform(X_all) # 用了全量数据")
print(" ✅ scaler.fit(X_train); scaler.transform(X_val) # 只用训练集")
print(" → 全量算 mean/std 会让测试集信息泄漏到训练")2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
八、自测题
Q1:消融实验中,去掉某个组件后性能下降了 2%。能直接说「这个组件贡献了 2%」吗?
答案
不能。至少要排除四种替代解释:
① 超参没重新调(最常见)
去掉组件后模型结构变了,最优 lr / weight_decay 可能也变了。用同一个 lr 测,性能下降可能只是「lr 不合适」。
必须:每个消融配置独立调超参。
② 差异在噪声范围内
2% 的下降,测试集多大?如果 SE 是 1.5%,2% 只有 1.3 个 SE,可能不显著。
必须:多 seed 报 ±std。
③ 训练不够
去掉组件后收敛更慢。需要给所有配置同样的训练轮数(或都训到收敛)。
④ 是间接作用
那个组件可能主要通过「影响其他组件」发挥作用。组合消融能检验这一点。
正确说法:「在固定超参和相同训练预算下,去掉组件 A 导致验证性能下降 2%(3 个 seed,std=0.3%)」。
加了限定条件才可信。
Q2:我的改进在验证集上提升了 1%,但测试集上没提升。最可能的原因是什么?
答案
最可能的原因:验证集被过度拟合(信息泄漏到决策过程)。
机制:你试了 20 组超参,每组都看验证集表现,最后挑最好的那个。这个「最好」包含了运气成分——在验证集上的优势部分是过拟合验证集。
这和第 2 篇的过拟合完全同构,只是过拟合的对象从「训练数据」变成了「验证集」。
其他可能:
- 测试集分布和验证集不同 —— 检查数据划分是否合理
- 测试集太小 —— SE 太大,1% 差异看不出来
- 改进是数据特定的 —— 需要在更多数据集上验证
- 真的没有改进 —— 验证集的 1% 是噪声
诊断方法:
1. 用多个 seed 跑,看验证集提升是否稳定
2. 如果 std > 1%,那验证集的提升就是噪声
3. 检查是否试了太多组超参(>10组就要警惕)2
3
根本的解法:用交叉验证(K 折),这样每个样本都当过验证集,能得到更可靠的估计。
一个实用建议:把测试集分成两半,一半用来做最终选择,一半做最终报告。这样能检测出「我是不是在验证集上过拟合了」。
Q3:一个论文报告「我们的方法在 ImageNet 上达到 85.0%,超过 ResNet-50 的 84.5%」。你怎么判断这个改进是否可信?
答案
按第 3.1 节的五步逐一检查:
① 差异是否显著? ImageNet val 有 50000 张图。
SE=500000.85×0.15≈0.16%
0.5% 的差异约 3 个 SE,边缘显著。需要知道他们是否做了配对检验。
② 是否多 seed? ResNet 这类成熟方法通常多次训练取平均。要看具体说明。
③ 超参是否公平? ★ 这是最关键的
ResNet-50 在 ImageNet 上被无数人调过超参(几十年积累),论文里用的可能是某篇经典实现的默认设置。作者的新方法呢?
如果新方法调了 50 组超参、ResNet 用了默认值,这个对比无效。
④ 消融实验完整吗? 看有没有证明「每个组件都有用」。
⑤ 代价是什么? 参数量、FLOPs、推理延迟。如果新方法精度高 0.5% 但慢 2 倍,那「更好」要看场景。
综合判断:
这个 0.5% 的差异在统计上勉强显著,但很可能是超参调优的产物。除非作者能提供:
- 在 2-3 个数据集上都有提升(不只是 ImageNet)
- 明确说明所有方法的调参协议
- 报告多 seed 结果
实践建议:自己复现。用同样的数据增强、训练轮数,重新训一遍 ResNet-50(它很快,1-2 天),看能不能达到 84.5%。如果你的 ResNet 能到 85.5%,那这 0.5% 的提升完全在调参范围内。
全套材料完成
回到 目录
Part 5 · 附录
不是读物,是工具。复习和面试前翻这里。
| 文档 | 内容 |
|---|---|
| 公式速查表 | 全部关键公式一页汇总,适合速记 |
| 面试高频问题 | 算法岗会问的 30 个问题,按主题分组,附答案要点 |
| llama_from_scratch.py | 从零实现的 LLaMA,163 行,可直接运行 |
| 离线单页版手册 | 全部内容打包成一个 HTML 文件,无网也能看 |
面试怎么用这两份
不要按顺序背。按题目反查:
能讲清楚「这个方法什么时候会失效」,比背下公式更管用。
离线单页版
深度学习进阶手册.html 是全部内容的单文件打包版,公式和代码高亮都已渲染。 下载后双击就能看,不需要网络,也不需要跑这个站点。
公式速查表
所有关键公式一页汇总。适合复习和面试前速记。
优化
梯度下降
θt+1=θt−η∇θL
SGD with Momentum
vt=βvt−1+gt,θt=θt−1−ηvt
展开:vt=gt+βgt−1+β2gt−2+⋯(指数加权平均)
RMSProp
st=βst−1+(1−β)gt2,θt=θt−1−st+ϵηgt
Adam
mtvtm^tv^tθt=β1mt−1+(1−β1)gt=β2vt−1+(1−β2)gt2=mt/(1−β1t)=vt/(1−β2t)=θt−1−ηv^t+ϵm^t
AdamW(解耦 weight decay)
θt=θt−1−ηv^t+ϵm^t−ηλθt−1
LLM 标准超参:lr=1e-4∼3e-4,β2=0.95,wd=0.1,clip=1.0,warmup 2000 步 + cosine decay
正则化
L2 / weight decay
θ←(1−ηλ)θ−η∂θ∂L
bias 和 1 维参数(BN/LN 的 γ,β)不做衰减。
L1
L~=L+λj∑∣θj∣
产生稀疏性(不可导点在 0)。
Dropout
训练时以概率 p 置零,并除以 1−p;推理时不丢弃。
反向传播
三条基本规则
| 运算 | 梯度 |
|---|---|
| z=a+b | ∂L/∂a=∂L/∂z |
| z=ab | ∂L/∂a=(∂L/∂z)b |
| z=a⊤W | ∂L/∂a=W(∂L/∂z)⊤,∂L/∂W=a(∂L/∂z)⊤ |
梯度消失/爆炸的量级
∂x1∂L∝l=1∏Ln1=n−L/2
sigmoid:σ′(z)≤0.25,每层乘 0.25
初始化
| 方法 | σw2 | 适用 |
|---|---|---|
| Xavier (Glorot) | nin+nout2 | tanh / sigmoid |
| He (Kaiming) | nin2 | ReLU / GELU |
| LLaMA | N(0,0.022) | 配合 RMSNorm |
归一化
BatchNorm
x^=σB2+ϵx−μB,y=γx^+β
统计维度:batch + 空间 → 训练/推理行为不同
LayerNorm
x^=σ2+ϵx−μ
统计维度:特征 → 与 batch 无关
RMSNorm
RMSNorm(x)=d1∑jxj2+ϵx⊙γ
省掉减均值,只有 γ。
CNN
卷积参数量
params=k×k×Cin×Cout+Cout
感受野
rl=rl−1+(kl−1)i=1∏l−1si,jl=jl−1+(kl−1)i=1∏l−1si
stride=1 简化:rl=1+∑i=1l(ki−1)
空洞卷积
有效感受野=k+(k−1)(d−1),需padding=d
ResNet
残差连接
y=F(x)+x,∂x∂y=∂x∂F+I
梯度连乘展开:
l=1∏L(I+Jl)=I+∑Jl+l<k∑JlJk+⋯
展开后第一项是 I,所以梯度不衰减。
Attention
Scaled Dot-Product Attention
Attention(Q,K,V)=softmax(dkQK⊤+M)V
为什么除 dk:点积方差 Var[q⋅k]=dk,std =dk。不缩放会让 softmax 饱和(实测 dk=64 时 max_p = 0.9990,缩放后 0.0350)。
形状流转:[n, d_k] × [d_k, m] → [n, m] → @ [m, d_v] → [n, d_v]
多头
MultiHead(Q,K,V)=Concat(head1,…,headh)WO
参数量和单头相同:4dmodel2
RoPE
θi=base−2i/d,Rmq=(cosmθsinmθ−sinmθcosmθ)q
核心性质:
⟨Rmq,Rnk⟩=⟨q,Rn−mk⟩=f(q,k,n−m)
内积只依赖相对位置。 base=10000(几何频率)。
现代 LLM 组件
SwiGLU
SwiGLU(x)=SiLU(xW1)⊗(xW3),SiLU(x)=xσ(x)
三个矩阵,中间维度取 38d 以保持参数量。
GQA
MHA: Q=[h,d] K=[h,d] V=[h,d]
GQA: Q=[h,d] K=[g,d] V=[g,d] ★ h/g = 4~8 最优
MQA: Q=[h,d] K=[1,d] V=[1,d]2
3
KV Cache
KV cache=2×nkv×L×dhead×nlayers×bytes
LLaMA-7B, L=4096, batch=2, fp16:
| 类型 | KV cache |
|---|---|
| MHA(32) | 2.15 GB |
| GQA(8) | 0.54 GB |
| MQA(1) | 0.07 GB |
缩放定律
Kaplan:L(N)=(NNc)αN,αN≈0.076
Chinchilla:Nopt∝C0.55,Dopt∝C0.45
Token/参数比最优 ≈ 20(但实践上为推理效率会超过这个值)。
损失函数
| 损失 | 来源假设 | 公式 | 任务 |
|---|---|---|---|
| MSE | 高斯 | N1∑(y−y^)2 | 回归 |
| MAE | 拉普拉斯 | $\frac{1}{N}\sum | y-\hat |
| CrossEntropy | 类别 | −logp^y | 多分类 |
| BCE | 伯努利 | 交叉熵(二分类) | 二分类 |
| NLL | — | −∑logp | 配合 LogSoftmax |
| InfoNCE | 对比 | −log∑jesim(i,j)esim(i,i) | 对比学习 |
关键:CrossEntropyLoss / BCEWithLogitsLoss 接受 logits,不是概率。
统计
偏差-方差分解
E[(y−y^)2]=Bias2+Var+σ2
标准误与置信区间
SE=np(1−p),95%CI=p±1.96SE
泛化差距
gap=train_loss−val_loss(或 train_acc−val_acc)
数值稳定
Loss Scaling
实际梯度=S⋅真实梯度
bf16 不需要(指数位与 fp32 相同,动态范围 10±38)。
梯度裁剪
g^=g⋅∥g∥max_norm当 ∥g∥>max_norm
等比缩放(不是逐元素裁剪,后者改变梯度方向)。
数值速查
| 概念 | 数值 |
|---|---|
| LLaMA-7B 参数量 | 6.7B(transformer 部分 5.37B) |
| LLaMA-7B 隐藏维度 | 4096 |
| LLaMA-7B 头数 / dk | 32 / 128 |
| ResNet-50 参数量 | 25.6M |
| FashionMNIST 模型参数量 | 669,706 |
| 输入 224² 全连接 vs 卷积 | 4.83e11 vs 1792(2.7亿倍) |
面试高频问题
算法岗面试会问的 30 个问题,按主题分组,附答案要点。按题目反查对应文档能深入理解。
一、基础概念(1-6)
1. 什么是过拟合?怎么判断?怎么解决? → Part1 第2篇
- 判断:训练指标好但验证指标差;两者差距(泛化差距)大
- 本质:模型容量 vs 数据量不匹配
- 解决(按性价比排序):加数据 > 数据增强 > 正则化 > 减容量
- 陷阱:验证误差低于训练误差 = 数据泄漏
2. L1 和 L2 正则的区别?为什么 L1 能产生稀疏? → Part1 第3篇
- L2 惩罚 λ∥θ∥2,梯度 λθ → 参数按比例收缩,不会恰好为 0
- L1 惩罚 λ∥θ∥1,不可导点在 0 → 参数被压到恰好为 0
- 深层原因:L1 对应拉普拉斯先验(密度在 0 处有尖峰),L2 对应高斯先验
- 联系:w←(1−ηλ)w−η∇L(L2 = 权重衰减)
3. 为什么 weight decay 不加在 bias 和 LayerNorm 上? → Part1 第3篇
- weight decay 惩罚的是「大权重 = 对输入敏感 = 不平滑」。bias 的作用是平移决策边界,对敏感度无贡献
- bias 初始化为 0,持续衰减会让它永远接近 0,等于关闭
- 标准写法:
if p.ndim <= 1 or name.endswith('.bias'): 不衰减
4. 交叉熵损失是怎么来的?为什么不用 MSE 做分类? → Part1 第1篇
- 来源:从「最大化数据似然」推导 → KL 散度 → 负对数似然 → 交叉熵
- 不用 MSE:MSE 对「概率应该多少」的约束弱(0.9 和 0.99 惩罚差距小),且梯度量级不合理
- 关键:CrossEntropyLoss 接收 logits 不是概率(内部做 softmax + log)
5. 什么是类别不平衡?怎么处理? → Part1 第2篇
- 表现:训练准确率很高但「全预测多数类」
- 处理:重采样(过采样少数/欠采样多数)、class weight、focal loss、调整决策阈值
- 诊断:看混淆矩阵,别只看 accuracy
6. Batch Norm 和 Layer Norm 的区别?为什么 LLM 用后者? → Part2 第6篇
| BatchNorm | LayerNorm | |
|---|---|---|
| 统计维度 | batch + 空间 | 特征 |
| 依赖 batch | 是 | 否 |
| 训练/推理行为 | 不同 | 相同 |
LLM 用 LN 的三个原因:batch 小(统计量不可靠)、变长序列(padding 污染)、推理需要确定性。
二、优化(7-13)
7. SGD 和 Adam 的区别?AdamW 又改进了什么? → Part1 第4篇
- SGD:只有 lr,无自适应
- Adam:一阶矩(动量)+ 二阶矩(自适应步长)
- AdamW:解耦 weight decay。Adam 里衰减会进入 mt,vt,被 vt1 缩放,行为不可预测;AdamW 直接 θ←θ−ηλθ
- LLM 为什么全用 AdamW:大数据下过拟合风险低,AdamW 更稳 + 可解耦其他衰减策略
8. 为什么 Adam 需要偏差修正? → Part1 第4篇
- m0=0 → 第一步 m1=(1−β)g1=0.1g1,严重低估
- 修正:除以 1−βt(几何级数的归一化因子)
- 效果:第一步的有效步长精确等于名义 lr
9. 为什么 LLM 用 β2=0.95 而不是 0.999? → Part1 第4篇
- 1−β21 = 二阶矩的有效窗口长度
- 0.999 → 1000 步窗口,训练中梯度尺度快速变化时会严重滞后
- 0.95 → 20 步窗口,能跟上当前状态
- 这是 HuggingFace / LLaMA 的默认配置,已成为事实标准
10. 为什么需要 warmup? → Part1 第4篇
- 初期梯度方向噪声大,直接用满 lr 会破坏已学到的结构
- Adam 早期 mt,vt 估计不准
- 配合 weight decay:实际衰减率 = ηλ,warmup 期间 lr 小 → 衰减也小(这是好事)
11. 梯度裁剪怎么做?为什么不用逐元素裁剪? → Part2 第5篇
- 等比缩放:
$\hat g = g \cdot \frac{c}{\|g\|}$,保持梯度方向不变 - 逐元素裁剪会改变方向,破坏「沿最陡下降方向」的前提
- LLM 标准配置:
clip_grad_norm_(model.parameters(), 1.0)
12. 为什么混合精度要用 loss scaling?bf16 呢? → Part2 第9篇
- fp16 范围 10±5,深层梯度 10−8 会下溢成 0
- loss scaling:梯度放大 S=65536 倍,
scaler.step()内部除回去,动态调整 S(遇 inf/nan 就减小) - bf16 指数位与 fp32 相同(10±38),不下溢,不需要 scaling
- 注意:用梯度裁剪时必须先
scaler.unscale_(optimizer)
13. 学习率怎么调?太大太小分别什么现象? → Part1 第4篇、Part1 第6章
| lr 过大 | lr 过小 |
|---|---|
| loss 震荡 / 变 nan | loss 几乎不降 |
| 训练像在发散 | 收敛极慢 |
- 最优 lr 随 batch size 缩放(线性缩放法则)
- LLM 微调:
1e-4 ~ 3e-4(AdamW);预训练:1e-4左右
三、深度学习原理(14-20)
14. 反向传播的链式法则怎么用? → Part2 第5篇
三条规则:加法→等值分发;乘法→交叉相乘;矩阵乘→外积。
∂W∂L=a⊤⋅∂z∂L
踩坑点:忘了 batch 平均的系数、矩阵乘法顺序写反(PyTorch 会静默广播出错结果)。
15. 梯度消失的本质是什么?残差连接为什么有效? → Part2 第8篇
- 本质:∂x1∂L∝∏ln1=n−L/2,深层指数衰减
- 残差的保证:∏l(I+Jl) 展开后第一项是 I,即使所有 Jl=0,梯度也原样传下去
- 实测:60 层无残差梯度 = 0,有残差 = 1.01e+05
16. 为什么不能用 0 初始化权重? → Part2 第9篇
- 全 0 → 输出 0 → 梯度 0 → 完全不训练
- 对称性问题:所有神经元梯度相同 → 永远同步更新 → 等效于一个神经元
- 例外:bias 可以初始化为 0;残差块最后一层的 BN 可以设 γ=0(让块初始为恒等)
17. Xavier 和 He 的区别?为什么 ReLU 网络要用 He? → Part2 第9篇
- Xavier:σ2=nin+nout2,兼顾前向和反向方差保持,适合对称激活
- He:σ2=nin2,补偿 ReLU 砍掉一半
- 实测(20 层):Xavier 衰减 5.87e5 倍,He 衰减 2.62 倍,默认初始化衰减 4.65e14 倍
18. Pre-LN 和 Post-LN 的区别?为什么现代 LLM 都用 Pre-LN? → Part3 第12篇
- Pre-LN:
x + Sublayer(LN(x))→ 残差通路是干净的恒等映射 - Post-LN:
LN(x + Sublayer(x))→ LN 夹在通路中间 - Pre-LN 需要 final norm(因为最后一个 block 的输出没过 norm)—— 这是最常见的实现 bug
19. RMSNorm 比 LayerNorm 少了什么?为什么效果不掉? → Part2 第6篇
- 少了:求均值、减均值、β 参数
- 核心:归一化的主要作用是「控制尺度」,不是「控制中心」
- 实测:RMSNorm 保留输入偏移(+0.426),但标准差一样被控制(0.905)
- 收益在算子融合,不在参数(省 d 个参数微不足道)
20. 什么是「退化问题」?和过拟合有什么区别? → Part2 第8篇
- 退化:网络加深,训练误差也上升 → 是优化问题,不是过拟合
- 过拟合:训练误差低,验证误差高
- 原因:优化器难以在深层网络中「找到」恒等映射的解
四、Transformer(21-27)
21. Attention 的公式?为什么除以 dk? → Part3 第10篇
Attention=softmax(dkQK⊤+M)V
为什么:Var[q⋅k]=dk → std = dk。不缩放 → logits std 太大 → softmax 饱和。
实测:dk=64 时不缩放 max_p = 0.9990(完全 one-hot,梯度消失),缩放后 0.0350。
22. 多头注意力的参数量比单头多吗? → Part3 第11篇
一样多,都是 4dmodel2。多头是把一个大注意力拆成 h 个小的,不是增加参数。
「多」的体现:h 个独立的注意力图,能学到不同的关系模式。
23. RoPE 为什么比绝对位置编码好? → Part3 第11篇
核心性质:⟨Rmq,Rnk⟩=f(q,k,n−m),内积只依赖相对位置。
四个优势:内置相对位置关系 / 不占表示维度 / 外推更好 / 是正交变换(不破坏语义结构)。
实测:同一 (q,k) 放不同位置,内积波动仅 1e-6。
24. 什么是 KV Cache?为什么能加速?代价是什么? → Part3 第13篇
- 原理:causal mask 保证已生成 token 的 kj,vj 永不改变 → 存下来复用
- 收益:计算量 O(n3)→O(n2)
- 代价:显存。LLaMA-7B, 4096 token, batch=2 → 2.15 GB
- 关键认知:Decode 是显存带宽瓶颈,注意力不是主要成本(MLP 占大头)
25. MHA / GQA / MQA 的区别?怎么选? → Part3 第13篇
| KV 头数 | KV cache(MHA 的比例) | |
|---|---|---|
| MHA | h | 100% |
| GQA | g | g/h |
| MQA | 1 | 3% |
论文结论:h/g=4∼8 时质量接近 MHA。实践:LLaMA3-70B 用 g=8,7B 用 g=8(h/g=4)。
26. Prefill 和 Decode 的瓶颈有什么不同?优化手段? → Part3 第13篇
| Prefill | Decode | |
|---|---|---|
| 瓶颈 | 算力 | 显存带宽 |
| 矩阵形状 | [2048,d]×[d,d] | [1,d]×[d,d] |
| 优化 | Flash Attention | 量化、GQA、continuous batching |
推论:Decode 带宽瓶颈 → 同权重下 batch 越大吞吐越高(这就是 continuous batching 的依据)。
27. 从头写一个 LLaMA Block 需要哪些组件? → Part3 第12篇
RMSNorm → Attention(SwiGLU 无) → 残差
RMSNorm → SwiGLU → 残差2
必备组件:RMSNorm、MultiHeadAttention(含 RoPE)、SwiGLU、causal mask、final norm。
参数量配置:4 个 d2(QKVO)+ 3×d×hffn(SwiGLU,hffn≈38d)。
五、研究方法(28-30)
28. 论文报告的提升 1%,可信吗?怎么验证? → Part4 第16篇
五步验证:
- 算置信区间(n=1000 时 1% 差异可能只有 1.4 个 SE)
- 3+ 随机种子报 ±std
- 配对检验(McNemar),不是独立 t 检验
- 给 baseline 也调参(最容易出问题的地方)
- 换数据集验证 + 报告代价
29. 消融实验怎么做才可信? → Part4 第16篇
四个原则:
- 一次只改一个东西
- 去掉后要重新调超参(最常被忽略)
- 多 seed
- 要做组合消融(检测交互作用)
组合消融的例子:单独去 A 掉 0.5%,单独去 B 掉 2%,但同时去掉掉 8% → A 和 B 强耦合,「B 重要 A 不重要」是错误结论。
30. 训练时 loss 变成 nan,怎么排查? → Part2 第9篇
按概率排序的原因:
| 原因 | 概率 | 解法 |
|---|---|---|
| 学习率太大 | 60% | 调小 |
| fp16 溢出 | 20% | bf16 或 loss scaling |
| log(0) | 10% | 检查 loss 实现 |
| 数据有 nan/inf | 8% | 清洗 |
| 梯度爆炸 | 2% | 梯度裁剪 |
定位工具:torch.autograd.detect_anomaly()
附录:高频「陷阱题」
Q:训练时 model.eval() 了会怎样? A:Dropout/BatchNorm 行为错误。eval 时 BN 用 running 统计量,train 时用当前 batch 统计量 → 结果不可复现。eval 不是「冻结」,是「切换行为」。
Q:为什么测试集不能用 shuffle=True? A:评估必须可复现,否则准确率的变化无法归因。而且训练集必须 shuffle(防过拟合),测试集绝不能。
Q:Batch Size 是不是越大越好? A:不是。大 batch → 梯度更准,但(a)泛化可能变差(小 batch 的噪声有正则化作用);(b)lr 要按比例放大;(c)显存受限。
Q:torch.optim.Adam 和 AdamW 我该用哪个? A:永远用 AdamW。Adam + weight_decay 在长训练里衰减行为不可预测,且和 AMP/分布式兼容性更差。
Q:为什么 LLM 不用 Dropout? A:预训练数据量远超参数量,过拟合风险低;Dropout 引入梯度噪声降低训练效率。但 LoRA 微调时反而要用(数据少)。
Q:模型 train loss 很低但效果不好,怎么排查? A:① 训练/验证数据集是否同分布 ② 验证集是否有泄漏 ③ 是否只监控了 loss 没监控业务指标 ④ 人工看 50 个错例(最有效的排查手段)
回到 目录