深度学习进阶从原理到研究方法 · 离线手册
总览与学习路线
Part 1 · 机器学习基础
Part 1 · 机器学习基础 01 机器学习的本质 02 泛化与偏差方差 03 正则化与容量控制 04 优化器 SGD→AdamW
Part 2 · 深度学习原理
Part 2 · 深度学习原理 05 反向传播推导 06 BatchNorm → RMSNorm 07 卷积与感受野 08 残差连接与架构 09 初始化与数值稳定
Part 3 · Transformer 与现代 LLM
Part 3 · Transformer 与现代 LLM 10 Attention 推导 11 多头与 RoPE 12 从 Transformer 到 LLaMA 13 KV Cache 与 GQA 14 缩放定律与 FlashAttention
Part 4 · 研究方法论
Part 4 · 研究方法论 15 怎么读论文 16 怎么设计实验
Part 5 · 附录
Part 5 · 附录 17 公式速查表 18 面试高频问题

深度学习进阶 · 从原理到研究方法

这套材料的定位:把「会写代码」变成「能判断」。

上一套 PyTorch 入门教你怎么做;这一套告诉你为什么这样做、什么时候不该这样做。
目标是让你能独立读论文、独立判断一个方案对不对、自己找出改进点。


这套材料和网上教程的区别

网上(包括大部分系统课)的问题是:告诉你「用残差连接能解决退化」,但不告诉你为什么能、什么条件下会失效、怎么判断一个方案是真的有效还是在过拟合。

这套材料的写法是:

  1. 推导优先。每个概念都用数学推导讲一遍,不用「可以理解为」糊弄。
  2. 给出边界。每个方法都说明它在什么条件下成立、什么时候崩。
  3. 可验证。每篇有可运行的代码,验证公式确实成立。

数学要求:会用导数、会矩阵乘法、理解条件概率即可。深度学习里用到的数学比你想的浅——真正的难点在直觉和判断力,不在公式难度。


目录与学习顺序

左侧目录树就是完整结构,这里再给一份可以一次看完的总表。

Part 1 · 机器学习基础(4 篇)

这一部分是地基。跳过它,后面所有内容都是空中楼阁。算法岗笔试和论文里反复出现的概念都在这里。

#文档核心问题
1机器学习的本质是什么学习=搜索?损失函数到底在衡量什么?
2泛化、过拟合与偏差方差训练集好测试集差,为什么?
3正则化与容量控制L1/L2 到底在做什么?为什么 weight decay 有时不等价?
4优化器:从 SGD 到 AdamWMomentum/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 年的主流设计。

#文档核心问题
10Attention 的数学推导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. 做最后的「边界条件」自测题
1
2
3
4
5

第 4 步是关键。如果你能讲清楚一个方法在什么情况下会失效,你才是真懂了。


关于代码

完整的 LLaMA 实现:llama_from_scratch.py,163 行,可运行:

bash
pip install torch
python llama_from_scratch.py    # 前向 + loss + 反向,全程通过
1
2

它包含 RMSNorm、RoPE、GQA、SwiGLU、Pre-LN 全部现代组件,可以直接作为你的实现参考。

文档里每段代码的输出都对应某个公式或结论(都是实测的)。如果你推出来的数和代码算出来的不一样,先信代码。


从第 1 篇开始 → 机器学习的本质是什么

Part 1 · 机器学习基础

这一部分是地基。跳过它,后面所有内容都是空中楼阁 —— 算法岗笔试和论文里反复出现的概念都在这里。

四条线,按顺序看:

#章节核心问题
01机器学习的本质是什么学习 = 搜索?损失函数到底在衡量什么?
02泛化、过拟合与偏差方差训练集好测试集差,为什么?
03正则化与容量控制L1/L2 到底在做什么?为什么 weight decay 有时不等价?
04优化器:从 SGD 到 AdamWMomentum / Adam 的更新量怎么推出来的?AdamW 修好了什么?

这一部分要建立的判断力

  • 看到「模型效果好」时,先问:好在哪个集合上、和什么比
  • 看到「加了正则化」时,先问:它约束的是什么,是参数范数还是函数复杂度
  • 看到「换了个优化器」时,先问:它改的是步长还是方向

顺序不能跳

第 2 篇(偏差方差)是后面所有内容的公共语言。 第 5 篇讲梯度消失时会直接引用它,第 14 篇讲缩放定律也会。

1 · 机器学习的本质是什么

核心问题:模型在学什么?损失函数在衡量什么?为什么「最小化损失」等于「学得好」?


一、一个不准确的直觉

多数入门材料会说:「机器学习就是让计算机从数据中学习规律」。这句话没错,但没有任何信息量。

更有用的定义是:

机器学习 = 在一个函数空间里搜索,使得在某个目标函数上取值最小。

三个关键词,逐个拆。

关键词一:函数空间

你写的模型 f(x; θ) 里,θ 是参数。参数固定,模型就固定了。

假设模型是线性回归 f(x) = wx + b,那么 w 和 b 就是两个自由变量。它们构成的二维平面就是这个模型的函数空间。

假设函数空间
线性回归(2 参数)二维平面
线性回归(10 参数)十维空间
一个 1 亿参数的 Transformer一亿维空间

机器学习的第一个关键事实:参数空间极其巨大,找到一个"好"的点,和找到一个"最优"的点,是两回事。

深度学习的全部困难都源于此——在一亿维空间里找最优解,实际做不到,只能不断改善。

关键词二:搜索

有了函数空间,接下来要决定「往哪个方向走」。这一步叫优化,由优化器(optimizer)完成。

在深度学习里,你写的这一行:

python
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-4)
1

就是决定了搜索策略。而 loss.backward() 是在计算方向(往哪走能让损失变小)。

「方向 + 步长」构成了优化的全部。 后面所有优化器的区别,本质上只是这两件事的不同做法。

关键词三:目标函数

你优化的是 loss_fn(pred, y)。这个函数衡量的东西,决定了模型学到什么。

  • 用 MSELoss → 模型学「预测值和真实值的平方距离」
  • 用 CrossEntropyLoss → 模型学「预测类别的对数似然」
  • 用对比学习损失(如 InfoNCE)→ 模型学「什么样本该靠近、什么该远离」

同一个模型架构,换个损失函数就是在做完全不同的事。 这是很多人换 loss 就想提升准确率却没效果的原因——他们换了,但没换对目标。


二、损失函数到底是什么

这一节是本文的核心。搞清楚损失函数的来源,后面所有选择都有依据。

从概率建模推导出来

假设你要做二分类。数据是 (x,y)(x, y)(x,y),其中 y∈{0,1}y \in \{0, 1\}y∈{0,1}。

建模思路:假设存在一个真实概率 p(y=1∣x)p(y=1 \mid x)p(y=1∣x),我们的模型输出一个 p^\hat{p}p^​。目标是让模型输出的概率尽可能接近真实概率。

用什么衡量两个概率分布的差异? KL 散度:

DKL(p∥q)=∑ipilog⁡piqiD_{KL}(p \| q) = \sum_i p_i \log\frac{p_i}{q_i} DKL​(p∥q)=i∑​pi​logqi​pi​​

直观含义:如果真实分布是 ppp,我们用 qqq 去编码,编码平均长度比最优编码多多少比特。越接近 0 越好。

(推论:为什么 KL 散度不对称——DKL(p∥q)≠DKL(q∥p)D_{KL}(p\|q) \ne D_{KL}(q\|p)DKL​(p∥q)=DKL​(q∥p)。这在后面的对比学习里会再次出现。)

从最大似然到交叉熵

KL 散度需要真实分布 ppp,但我们不知道。改用「最大似然」:找到一组参数,让观测到的数据出现的概率最大。

L=−∑ilog⁡p^(yi∣xi)L = -\sum_i \log \hat{p}(y_i \mid x_i) L=−i∑​logp^​(yi​∣xi​)

即负对数似然(NLL)。因为最大化 ∏p\prod p∏p 等价于最小化 ∑−log⁡p\sum -\log p∑−logp。

展开二分类(y∈{0,1}y \in \{0,1\}y∈{0,1}):

L=−1N∑i[yilog⁡p^i+(1−yi)log⁡(1−p^i))L = -\frac{1}{N}\sum_i \left[y_i \log \hat{p}_i + (1-y_i)\log(1-\hat{p}_i)\right) L=−N1​i∑​[yi​logp^​i​+(1−yi​)log(1−p^​i​))

这就是交叉熵损失(Cross-Entropy Loss)。

所以:交叉熵损失不是拍脑袋设计的,它就是「最大化观测数据的似然」在分类任务上的形式。

多分类版本

模型输出 KKK 个 logits z1,…,zKz_1,\dots,z_Kz1​,…,zK​,用 softmax 转成概率:

p^k=ezk∑jezj\hat{p}_k = \frac{e^{z_k}}{\sum_j e^{z_j}} p^​k​=∑j​ezj​ezk​​

损失:

L=−1N∑ilog⁡p^i,yiL = -\frac{1}{N}\sum_i \log \hat{p}_{i,y_i} L=−N1​i∑​logp^​i,yi​​

关键细节:PyTorch 的 nn.CrossEntropyLoss() 接受 logits,不是概率。 它内部自己做 softmax。

python
# ✅ 正确:传 logits
loss = nn.CrossEntropyLoss()(model(x), y)

# ❌ 错误:先 softmax 了再传,模型会「过度自信」
loss = nn.CrossEntropyLoss()(torch.softmax(model(x), dim=1), y)
1
2
3
4
5

为什么传 logits 更好? 数值稳定性。logits 可以是任意实数,softmax 内部用了减最大值的技巧避免溢出;如果你先 softmax,log(0) 会产生 inf/nan。

回归任务

如果 yyy 是连续值,通常假设 y∼N(f(x),σ2)y \sim \mathcal{N}(f(x), \sigma^2)y∼N(f(x),σ2)(高斯假设),最大化似然后得到均方误差:

L=1N∑i(f(xi)−yi)2L = \frac{1}{N}\sum_i (f(x_i) - y_i)^2 L=N1​i∑​(f(xi​)−yi​)2

所以 MSE 也是从概率假设推出来的,不是随便定义的。不同分布假设 → 不同损失函数:

分布假设对应损失适用任务
高斯MSE(L2)回归
拉普拉斯MAE(L1)回归、抗 outliers
类别(softmax)CrossEntropy多分类
伯努利BCE二分类

「选 loss 就是选分布假设」 —— 这是我觉得最值得记住的一句话。

L1 vs L2(理解正则化的基础)

损失函数概率视角
MSE∣y−y^∣22|y - \hat{y}|_2^2∣y−y^​∣22​高斯噪声
MAE∣y−y^∣1|y - \hat{y}|_1∣y−y^​∣1​拉普拉斯噪声

两者对 outlier 的敏感度差异巨大:MSE 对大误差惩罚是平方级的,MAE 是线性的。

所以:数据干净用 MSE,有 outliers 用 MAE。

python
# PyTorch 里对应
nn.MSELoss()      # 平方
nn.L1Loss()       # 绝对值
nn.HuberLoss()    # 两者混合,小误差用平方、大误差用绝对值
1
2
3
4

三、EM 算法:一个必须懂的例子

EM 算法不属于深度学习,但它是最能说明「损失函数从哪来」的经典案例,而且是理解 VAE、GAN 的前提。

问题:有一堆数据来自两个不同均值的正态分布(一个 0 号月亮,一个 1 号月亮),但每个样本的标签被抹掉了。已知两个高斯的均值和方差,怎么估计每个样本属于哪个?

难点:分配依赖参数,参数依赖分配。死循环。

EM 的思路:

E 步(Expectation):用当前参数,估计每个样本"属于各类的概率"
                    这得到一个"软标签",而不是硬分配
M 步(Maximization):用这些软标签重新估计参数(本质是加权最小二乘)

重复 E → M → E → M ... 直到参数不再变化
1
2
3
4
5

为什么这个"绕"能用? 因为 E 步构造了一个下界(Jensen 不等式),每次迭代让下界上升,参数收敛到最大似然估计。

和深度学习的关系:

VAEGAN
隐变量有(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:损失函数的选择如何改变模型学到的「形状」

python
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}   ← 更接近正确值")
1
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.5
1
2

L2 的斜率被单个异常值拉偏了近 3 倍。这就是为什么回归任务里数据清洗和抗 outlier 损失这么重要。

实验 2:softmax 的数值稳定性

python
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))                # 正常
1
2
3
4
5
6
7
8
9
10
11

原理:softmax 有平移不变性 softmax(z)=softmax(z−max⁡z)\text{softmax}(z) = \text{softmax}(z - \max z)softmax(z)=softmax(z−maxz),减去最大值后所有指数都 ≤ 1,不会溢出。

这就是为什么 CrossEntropyLoss 要吃 logits 而不吃概率 —— 见前面「误解」部分的解释。

实验 3:Momentum 到底在做什么

python
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]])
1
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:过拟合的可视化

python
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))])
1
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:如果损失函数是 −log⁡p^(y∣x)-\log \hat{p}(y|x)−logp^​(y∣x),当模型给真实标签分配的概率是 0.001 时,损失大概是多少?为什么要这么设计?

答案

−log⁡(0.001)≈6.9-\log(0.001) \approx 6.9−log(0.001)≈6.9

为什么要这么惩罚:

  • 用「误差」衡量的话,给 0.9 和给 0.001,差 0.899 → 惩罚力度差不多
  • 用对数的话,差 6.9 → 相差一个数量级

这就带来了梯度的自然缩放:在概率已经很低时,模型每提升一点概率,减少的损失很多,梯度很大,模型会集中力气去修正最「离谱」的预测。这正是我们想要的。

代价是:过度自信的错误预测会被重罚(给正确标签 1e-6 的概率 → 损失 13.8)。这也是 label smoothing(给标签加一点平滑)存在的原因。

Q4:为什么说「EM 是联合优化的替代方案」?它在什么情况下会失效?

答案

EM 只适用于「联合似然有闭式解或易优化」的情况,且能保证单调收敛。

失效场景:

  1. 参数空间的似然不可分解(EM 要求可分解的联合分布)
  2. 初始化太差 → 收敛到局部最优(EM 只保证局部收敛到某个驻点,不保证全局最优)
  3. KL 散度取反方向(变分下界)时,对数似然是凹的才能保证收敛,否则不保证

实际意义:VAE 用的是「变分 EM」的思路,能训但会近似——因为 decoder 的似然往往不是高斯的。所以 VAE 的输出天然是模糊的(这也是它生成质量不如 GAN 的原因之一)。


下一篇 → 泛化、过拟合与偏差方差

上一级: 目录

2 · 泛化、过拟合与偏差方差

核心问题:为什么训练集上表现好,测试集上却不行?如何定量判断一个模型是真好还是只是记住了数据?


一、先纠正一个直觉

多数人第一次遇到「过拟合」时的反应是:「模型把训练数据记住了,所以测试就差」。

这个说法不准确,但有用。准确的表述是:

模型的假设空间包含了训练集的一个特例(完美拟合),而我们优化算法恰好找到了这个特例。测试数据不属于这个特例,所以失效。

关键点:过拟合不是模型的缺点,而是「模型容量 vs 数据量」不匹配的表现。 容量没有错,错的是容量相对于数据量太大了。

推论很重要:

  • 加数据能缓解过拟合(同一个高容量模型,数据越多越不容易找到「特例」)
  • 减小模型容量也能缓解(限制假设空间)
  • 改正则化方式也能缓解(但不是缩小空间,而是改变「哪些解更容易被优化找到」)

二、偏差-方差分解(Bias-Variance Decomposition)

这是本文最有价值的部分。它把「误差」拆成了三个可分别优化的部分。

2.1 误差的三项分解

对单个样本 xxx,模型的预测误差可以分解为:

E[(y−y^(x))2]=Bias2[y^(x)]⏟欠拟合+Var[y^(x)]⏟过拟合+σ2⏟噪声(不可消除)\mathbb{E}\big[(y - \hat{y}(x))^2\big] = \underbrace{\text{Bias}^2[\hat{y}(x)]}_{\text{欠拟合}} + \underbrace{\text{Var}[\hat{y}(x)]}_{\text{过拟合}} + \underbrace{\sigma^2}_{\text{噪声(不可消除)}} 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%,波动巨大 → 【高方差,可能低偏差】
1
2
3
4
5
6
7
8
9
10
11

核心:偏差和方差是同一个模型的两种不同失败模式,不是两个独立的可调参数。

  • 模型太弱:预测系统性地错 → 偏差大
  • 模型太强:预测不稳定地错(换个数据就变) → 方差大

2.3 U 型曲线:最重要的一张图

横轴 = 模型容量(或训练轮数),纵轴 = 误差:

误差
 ↑
 │  方差(过拟合)        ←── U型曲线 ──→      偏差(欠拟合)
 │      ╱                          ╲
 │     ╱    最佳点                    ╲
 │    ╱        ●                      ╲
 │   ╱                                  ╲
 │  ╱                                    ╲
 └──────────────────────────────────────────→ 模型容量 / 训练轮数
1
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:时序数据随机切分

如果数据有时间顺序,随机切分会泄漏未来信息。

❌ 错误:随机切分 → 训练集里有"未来"的数据,测试集有"过去"
✅ 正确:按时间切分(用最近的一段做测试)
1
2

实际项目里的判断标准很简单:你能拿到「上线后」的数据吗? 评估方案就该模拟那个场景。

坑 3:数据本身不独立(同一个人多张照片)

如果同一个人的多张照片散落在训练集和验证集,模型学到的是「认出这个人」而不是「认出这个类别」。

必须按「组」切分(GroupSplit),保证同组数据全在一侧。

K 折交叉验证

数据少的时候,val 划分不可靠(val 太小 → 指标方差大)。用 K 折:

数据分成 K 份,轮流当验证集,其余当训练集
→得到 K 个 (train_score, val_score)
→ 报告平均值 ± 标准差
1
2
3
python
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):
    ...
1
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: 逐项试正则化,量化每项的贡献(做消融实验)
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16

Step 3 是关键,很多「模型问题」实际上是数据问题。判断方法:人工看 50 个被分错的样本,如果里面有一半是标注错误,那问题在数据不在模型。


五、动手实验

实验 1:画出过拟合的 U 型曲线

这是本篇最重要的实验。用决策树深度作为「模型容量」的旋钮,扫描并观察训练/验证曲线的分叉。

(为什么用决策树而不是多项式回归?因为决策树在深度足够大时能做到训练集完美拟合(loss→0),U 型非常清晰。逻辑回归有 L2 正则压着,即使多项式阶数拉满也不会真正过拟合——我实测过,训练和验证曲线几乎重合,看不出教学效果。)

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

实测输出:

 深度   训练准确率   验证准确率      泛化差距
   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.100
1
2
3
4
5
6
7
8
9

你应该观察到:

  • 训练准确率单调上升,深度 13 之后达到 1.000(完全记住训练集)
  • 验证准确率在深度 6 达到峰值 0.920,之后掉头向下
  • 泛化差距从 -0.01 一路增长到 +0.10

实验 2:泛化差距告诉你「差多少」

python
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 模型在记噪声
1
2
3
4

差距是模型「过度自信」的程度。0.021 说明模型学到的基本都对;0.100 说明它花了大量容量去记住噪声。

注意深度 13~15 的训练准确率完全一样(都是 1.000),但深度 1~2 差距是负数——这是因为小模型训练得不够充分,还没开始过拟合。负差距不是 bug,是「欠拟合且还没到过拟合阶段」的正常现象。

实验 3:加数据能不能缓解过拟合

这是最有说服力的实验:固定模型(一个会完全记住训练集的决策树),只变数据量。

python
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}")
1
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.108
1
2
3
4
5

你会观察到:训练准确率恒定在 1.000(模型从头到尾都在过拟合),但验证准确率随数据量上升(0.867 → 0.906),差距整体在收窄。

这直接证明了开头的结论:过拟合的本质是「容量 vs 数据量」不匹配,加数据是最直接的解法。

注意:这个提升不是单调的(400→800 反而降了)。因为数据变少时验证集本身也在变小,指标方差变大。小数据上的指标波动是正常的,这也是为什么小数据要做 K 折交叉验证。

我原本用 MLP(64,64)跑这个实验,但 2000 条数据的月牙对 MLP 来说太简单了,训练准确率一直在 0.91~0.94 徘徊,根本没过拟合,也就看不到「训练集恒定、验证集上升」这个对比。做实验时如果观察不到预期现象,要先怀疑实验设置,而不是硬编一个解释。

实验 4:检测数据泄漏

python
# 检测泄漏的信号:验证误差 > 训练误差(超出正常波动范围)
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])
1
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:为什么「训练误差和验证误差差不多」不代表模型没问题?

答案

因为两个模型可能同时很差。有几种典型情况:

  1. 欠拟合:训练误差 55%、验证误差 54%。差距很小,但两个都很高——模型根本没学到东西
  2. 数据太简单/标签有噪声:训练验证都是 50%,但真实上限就 50%(数据本身随机标注)
  3. 指标选错:accuracy 都是 50%,但 A 类的 recall 是 100%、B 类是 0
  4. 类别不平衡:全预测多数类也能拿到高准确率

结论:差距小只说明「没有过拟合」,不说明「拟合得好」。 必须同时看绝对值和业务指标。这就是为什么规范的做法是同时监控训练和验证的多个指标。

Q3:Dropout 在训练和推理时行为不同。如果忘了在推理时调 model.eval(),会发生什么?为什么这个 bug 很难发现?

答案

推理时仍然随机丢弃神经元 → 相当于用一个「随机衰减的次优网络」做预测。

为什么难发现:

  1. 代码完全不会报错
  2. 准确率会下降但不会崩到离谱,可能只降几个百分点,很容易被归因为「随机波动」
  3. 同一个模型的验证指标每次跑都不一样 → 看起来「验证集有噪声」,反而会让人去怀疑数据划分

排查方法:多次评估同一模型,如果指标方差异常大(比如 ±3%),大概率是漏了 eval()。因为 eval() 缺失时每次前向都不同。

这也是为什么 model.eval() 和 torch.no_grad() 总是成对出现:前者管正确性,后者管效率。

Q4:一个模型在测试集上 85%,另一个 86%。这个差异有意义吗?

答案

大概率没有意义,除非你能算出置信区间。

在 10000 个测试样本上,85% 和 86% 的差别是 100 个样本。估计标准误约为:

SE≈p(1−p)n=0.855×0.14510000≈0.0035=0.35%SE \approx \sqrt{\frac{p(1-p)}{n}} = \sqrt{\frac{0.855 \times 0.145}{10000}} \approx 0.0035 = 0.35\% 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(θ)⏟复杂度惩罚\min_{\theta} \underbrace{L(\theta)}_{\text{数据拟合}} + \underbrace{\lambda \cdot R(\theta)}_{\text{复杂度惩罚}} θmin​数据拟合L(θ)​​+复杂度惩罚λ⋅R(θ)​​

区别只在R(θ)R(\theta)R(θ) 的定义:

方法R(θ)R(\theta)R(θ)效果
L2 正则 / weight decay∣θ∣22|\theta|_2^2∣θ∣22​参数整体变小,均匀收缩
L1 正则∣θ∣1|\theta|_1∣θ∣1​参数变稀疏(部分归零)
Dropout—(随机丢弃)制造集成效果
Early stopping—(早停)限制有效训练步数
数据增强—(扩充数据)增加有效数据量

统一理解:这些方法都在限制模型的有效容量,只是途径不同。

一个关键区分:限制「空间」还是限制「路径」

这两类正则化作用机制完全不同,这是本文最重要的观点:

限制假设空间限制优化路径
手段L1/L2、网络结构变小Dropout、Early stopping、BatchNorm、初始化
效果最优解本身变了最优解没变,但你到不了那里
类比房子盖小一点走小路绕过去

这个区分解释了为什么有些正则化会互相冲突。比如强L2(缩小空间)+ 强 Dropout(限制路径),可能两个一起用反而不如单独用。


二、L2 正则(权重衰减)

2.1 数学推导:为什么正则项能防止过拟合

目标函数:

L~(θ)=L(θ)+λ2∥θ∥22\tilde{L}(\theta) = L(\theta) + \frac{\lambda}{2}\|\theta\|_2^2 L~(θ)=L(θ)+2λ​∥θ∥22​

对参数求梯度:

∂L~∂θj=∂L∂θj+λθj\frac{\partial \tilde{L}}{\partial \theta_j} = \frac{\partial L}{\partial \theta_j} + \lambda \theta_j ∂θj​∂L~​=∂θj​∂L​+λθj​

梯度下降一步:

θj←θj−η(∂L∂θj+λθj)=(θj−η∂L∂θj)−ηλθj\theta_j \leftarrow \theta_j - \eta\left(\frac{\partial L}{\partial \theta_j} + \lambda \theta_j\right) = (\theta_j - \eta \frac{\partial L}{\partial \theta_j}) - \eta\lambda\theta_j θj​←θj​−η(∂θj​∂L​+λθj​)=(θj​−η∂θj​∂L​)−ηλθj​

整理一下:

θj←(1−ηλ)θj−η∂L∂θj\boxed{\theta_j \leftarrow (1 - \eta\lambda)\theta_j - \eta\frac{\partial L}{\partial \theta_j}} θj​←(1−ηλ)θj​−η∂θj​∂L​​

看到了什么? 参数更新 = 数据梯度的修正 − 参数自身的衰减。

每一步参数都会按比例 (1−ηλ)(1-\eta\lambda)(1−ηλ) 缩小。这就是「权重衰减」这个名字的来源——参数被持续地往零的方向拉。

2.2 拉普拉斯先验的解释(贝叶斯视角)

L2 正则等价于假设参数服从标准正态先验:

p(θ)∝exp⁡(−λ2∥θ∥2)p(\theta) \propto \exp\left(-\frac{\lambda}{2}\|\theta\|^2\right) p(θ)∝exp(−2λ​∥θ∥2)

由贝叶斯公式:

p(θ∣D)∝p(D∣θ)⏟似然(对应 L(θ))×p(θ)⏟先验(对应 λR(θ))p(\theta|D) \propto \underbrace{p(D|\theta)}_{\text{似然(对应 } L(\theta)\text{)}} \times \underbrace{p(\theta)}_{\text{先验(对应 } \lambda R(\theta)\text{)}} p(θ∣D)∝似然(对应 L(θ))p(D∣θ)​​×先验(对应 λR(θ))p(θ)​​

对应关系:

频率派(正则化)贝叶斯派(先验)
数据拟合项 L(θ)L(\theta)L(θ)似然 p(D∣θ)p(D|\theta)p(D∣θ)
正则项 λR(θ)\lambda R(\theta)λR(θ)参数先验 p(θ)p(\theta)p(θ)
λ\lambdaλ 大先验强(更相信先验)

同理:

正则对应先验参数效果
L2高斯 N(0,1/λ)N(0, 1/\lambda)N(0,1/λ)平滑收缩,不会恰好为 0
L1拉普拉斯,密度在 0 处有尖峰稀疏

L1 稀疏性的数学原因:拉普拉斯分布的 PDF 在 0 处最高,所以后验在 0 处的概率最大 → 很多参数恰好被压到 0。

2.3 L2 为什么让模型更平滑

一个直觉解释:L2 惩罚大参数。

回想第 1 篇——MSE 的目的是拟合噪声。而大参数意味着对输入敏感(y^=Wx\hat{y} = Wxy^​=Wx,WWW 越大,输入的小变化引起输出的大变化)。所以 L2 惩罚等价于惩罚「模型对输入的敏感度」,也就是在提升平滑性。

平滑性 = 模型对微小扰动不敏感 = 对噪声不敏感 = 泛化更好。


三、L1 正则与稀疏性

L~(θ)=L(θ)+λ∥θ∥1\tilde{L}(\theta) = L(\theta) + \lambda \|\theta\|_1 L~(θ)=L(θ)+λ∥θ∥1​

L1 的不可导点在 0,这导致它的解倾向于落在坐标轴上(很多参数恰好为 0)。

L1 vs L2 的本质差异

L1L2
惩罚函数∣x∣1|x|_1∣x∣1​,斜的x2x^2x2,光滑
稀疏性有,精确为 0无,只是变小
梯度子梯度:+1+1+1 或 −1-1−12x2x2x,与 x 成正比
收缩方式均匀收缩小参数,大参数收缩很少与参数大小成比例收缩
适合特征选择、稀疏模型通用默认

为什么 L1 产生稀疏:直觉上,L1 对大参数的惩罚力度和 L2 不同——

  • L2 惩罚 λx2\lambda x^2λx2:xxx 越大,惩罚越重(所以大参数被压得厉害)
  • L1 惩罚 λ∣x∣\lambda |x|λ∣x∣:惩罚恒定(所以大参数和小参数受到的「绝对压力」一样)

结果:大参数被 L2 压下去了但 L1 不管;小参数在 L1 下更容易被压到 0。L1 因此偏向于「要么留下一个大的,要么完全不要」。

什么时候用 L1

  • 特征选择:特征是几万个混杂在一起的信噪比时
  • 需要可解释模型:能说清哪些特征重要
  • 模型压缩:相比剪枝后直接得到小模型更方便

实践中深度学习里很少用 L1,因为 weight decay(L2)更稳定,且没有稀疏性需求。


四、weight decay vs L2 正则:什么时候不等价

这是最容易被忽略但面试常问的点。

4.1 理论上等价

L2 正则:把 λ∥θ∥2\lambda\|\theta\|^2λ∥θ∥2 加到 loss 上。

optimizer.zero_grad()
loss = loss_fn(pred, y) + lambda * sum(p.pow(2).sum() for p in model.parameters())  # ← 加到 loss
loss.backward()
1
2
3

weight decay:优化器内部直接衰减参数,不经过 loss。

python
optimizer = torch.optim.SGD(model.parameters(), lr, weight_decay=lambda)
1

4.2 两者的真实差异

L2 正则weight decay
实现手动加到 loss优化器内部做 p -= lr*wd*p
梯度可见性体现在 loss 数值里loss 里看不到
系数与 lr 的关系λ\lambdaλ 独立实际衰减率 = ηλ\eta\lambdaηλ,依赖 lr
梯度裁剪的交互会影响裁剪前的梯度裁剪后再衰减,行为不同
AMP / bf16 兼容手写可能有问题优化器原生支持

关键差异:weight decay 的实际强度和 lr 耦合。 因为衰减量是 ηλθ\eta\lambda\thetaηλθ——如果你把 lr 改大了,衰减也跟着变大。L2 正则则不受 lr 影响。

实践建议:优先用 weight decay(优化器原生支持,和 AMP/分布式兼容),不要手写 L2 加到 loss 上。

4.3 哪些参数不该衰减

这是现代实践里很重要的一点:bias 和 LayerNorm/BatchNorm 的参数不应该做 weight decay。

原因:

  • bias 的作用是平移决策边界,衰减它没有正则化意义
  • BN/LN 的 γ,β\gamma, \betaγ,β 是缩放平移参数,它们应该自由调整
python
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)
1
2
3
4
5
6
7
8
9
10
11
12
13
14

这是所有主流 LLM 微调代码的标准写法(如 HuggingFace、LLaMA 微调脚本)。面试如果问「你怎么调weight decay」,能说出这一段是加分项。


五、Dropout 的原理

5.1 做了什么

训练时,每个神经元以概率 ppp 被丢弃(置零)。测试时,不做丢弃,而是把激活值除以 1−p1-p1−p(PyTorch 用 inverted dropout:训练时除以 1−p1-p1−p)。

5.2 为什么有效:集成视角

一个关键洞察:Dropout 每次前向都在训练一个不同的子网络。

第1次前向: 用神经元 {1,3,5,7}
第2次前向: 用神经元 {2,4,6}
第3次前向: 用神经元 {1,2,6,8}
...
1
2
3
4

训练 NNN 次等于训练了 NNN 个不同结构的模型,每个都只看到部分数据。推理时用全部神经元,相当于对这些子网络的集成平均。

E[输出]≈各个子网络输出的加权平均\mathbb{E}[\text{输出}] \approx \text{各个子网络输出的加权平均} 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 正则的权重衰减效应

亲手看到参数被「拉向零」。

python
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}")
1
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.04749
1
2
3
4

观察:l2l_2l2​ 越大,∥w∥\|w\|∥w∥ 越小,训练损失越高。这就是「偏差-方差权衡」的实证——压得越狠,训练越差但泛化可能越好。

注意这个任务是无噪声线性回归(生成数据时用了同一个 torch.randn(10,1)),所以真实最优 λ=0\lambda = 0λ=0。如果你的任务有噪声,λ>0\lambda > 0λ>0 应该在验证集上真的更好——这是检验正则化是否有效的正确方式。

实验 2:L1 产生稀疏,L2 不产生

python
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()]}")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25

你会观察到:L1 产生大量恰好为 0 的参数(稀疏),L2 只是把所有参数缩小。

实验 3:dropout 的集成视角

验证「不同 dropout 模式 = 不同子网络」。

python
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))
1
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 下的衰减量。

python
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}")
1
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.00
1
2
3
4

三个结论:

  1. 公式精确成立:衰减量就是 ηλθ\eta\lambda\thetaηλθ,比值全是 1.00。
  2. 衰减量随 lr 线性变化:lr 放大 10 倍,衰减量也放大 10 倍。这正是 weight decay 和 L2 的差异所在。
  3. 如果你用 L2 加到 loss 上,衰减量会是 λθ\lambda\thetaλθ(不含 lr),和 lr 无关。

这个实验也解释了为什么「把 lr 调大一倍,正则化强度也跟着变强一倍」。用 lr schedule 时(比如 warmup),实际正则强度是在变的。

我最初试过在完整训练里对比不同 lr 的 ∥w∥\|w\|∥w∥,结果看不出差异——因为任务很快就收敛了,收敛后的 ∥w∥\|w\|∥w∥ 由数据拟合需求主导,衰减的影响被淹没。做实验时如果观察不到预期现象,先怀疑实验设置。 隔离变量(梯度置零)是更可靠的做法。


八、自测题

Q1:为什么 bias 不做 weight decay?

答案

三个理由:

  1. 无正则意义:weight decay 惩罚的是「模型对输入的敏感度」(大权重 → 对输入敏感 → 平滑性差)。而 bias 的作用是平移决策边界,对输入敏感度没有贡献。衰减 bias 只会让决策边界往原点靠,纯属有害。
  2. 数学上:bias 通常初始化为 0,如果持续衰减,它永远接近 0,等于关闭了 bias 项。
  3. 和 norm 层同理:LayerNorm 的 γ,β\gamma, \betaγ,β、BatchNorm 的 γ,β\gamma, \betaγ,β 都是 1 维参数(p.ndim == 1),它们的作用是缩放和平移,同样不该衰减。

所以标准做法是:if p.ndim <= 1 or name.endswith('.bias'): 不衰减。这是所有主流 LLM 微调代码的标准配置。

Q2:为什么 Dropout 在推理时要除以 1−p1-p1−p(或者反过来乘保留概率)?

答案

为了让训练和推理的期望输出一致。

训练时,神经元以概率 ppp 被置零。某个神经元的激活值 xxx:

  • 以概率 1−p1-p1−p 保留:期望贡献 (1−p)x(1-p)x(1−p)x
  • 以概率 ppp 丢弃:贡献 0

所以期望是 (1−p)x(1-p)x(1−p)x。如果不修正,训练时的期望输出会比推理时小 (1−p)(1-p)(1−p) 倍,推理时输出会系统性偏大。

两种修正方式:

  • inverted dropout(PyTorch 用):训练时激活值除以 1−p1-p1−p,推理时不修正
  • 训练时不修正,推理时所有激活值乘 1−p1-p1−p

两种等价,PyTorch 选前者是为了推理时更简单高效。

Q3:现代 LLM(如 LLaMA)为什么不用 Dropout?

答案

三个原因:

  1. 过拟合风险低:预训练数据量(万亿 token 级)远超参数量,模型几乎不会过拟合。这是数据规模化的直接收益。
  2. 梯度噪声影响效率:大规模预训练是吞吐敏感的。Dropout 引入随机梯度噪声,需要更多步数收敛,训练效率下降。
  3. 有更好的替代手段:
    • weight decay(AdamW)控制参数规模
    • 大量数据 + 更深的网络本身就有正则化效应
    • 现代预训练配方里正则化项占很小比重

但 LoRA 等高效微调方法下,情况不同:数据量小(几万条),过拟合风险高,所以微调时反而会用 dropout(LoRA 原论文就用 0.05)。这说明正则化的选择取决于数据量与模型容量的比值,不存在「永远不用」或「永远用」的规则。

Q4:手写 L2 正则加到 loss 上和用 weight_decay,如果你的 lr schedule 用了 warmup,两种会有区别吗?

答案

会有区别,而且值得注意。

warmup 期间 lr 很小(接近 0),此时:

  • weight decay 的衰减量 ηλθ\eta\lambda\thetaηλθ 也很小 → 几乎没有衰减
  • L2 加到 loss:梯度 λθ\lambda\thetaλθ 是常数项,不依赖 lr → 照常衰减

所以 warmup 初期,L2 正则的实际强度比 weight decay 更强。

如果 warmup 阶段 loss 出现异常(比如初期就过拟合),可以考虑:

python
# 让 weight decay 跳过 warmup(部分实现的做法)
if current_step < warmup_steps:
    for group in optimizer.param_groups:
        group['weight_decay'] = 0.0
1
2
3
4

不过实践中大多数框架的 weight decay 实现里,这两者差异被忽略了,因为 warmup 阶段的权重还很小,衰减多少都无所谓。


下一篇 → 优化器:从 SGD 到 AdamW

上一级: 目录 · 上一篇

4 · 优化器:从 SGD 到 AdamW

核心问题:Momentum、Adam 的更新量是怎么推出来的?AdamW 修好了Adam 的什么缺陷?为什么 LLM 全用 AdamW?


一、优化器要解决的真问题

朴素 SGD 的更新:θ←θ−ηg\theta \leftarrow \theta - \eta gθ←θ−ηg

三个致命缺陷:

  1. 步长无法自适应:某些参数梯度大、某些小,SGD 对它们的处理完全一样,导致收敛慢
  2. 梯度方向震荡:高维空间里梯度方向在不同维度上符号频繁变化,Z 字形前进
  3. 学习率必须手调:太大学习率发散,太小收敛慢,且不同参数需要不同 lr

所有优化器的改进都是围绕这三点。


二、SGD with Momentum

2.1 公式与推导

vt=βvt−1+gt,θt=θt−1−ηvt\boxed{v_t = \beta v_{t-1} + g_t, \qquad \theta_t = \theta_{t-1} - \eta v_t} vt​=βvt−1​+gt​,θt​=θt−1​−ηvt​​

其中 vtv_tvt​ 是速度(velocity),v0=0v_0 = 0v0​=0。

展开这个递推式:

vt=βvt−1+gt=β(βvt−2+gt−1)+gt=β2vt−2+βgt−1+gtv_t = \beta v_{t-1} + g_t = \beta(\beta v_{t-2} + g_{t-1}) + g_t = \beta^2 v_{t-2} + \beta g_{t-1} + g_t vt​=βvt−1​+gt​=β(βvt−2​+gt−1​)+gt​=β2vt−2​+βgt−1​+gt​

继续展开到最初:

vt=gt+βgt−1+β2gt−2+⋯+βtv0v_t = g_t + \beta g_{t-1} + \beta^2 g_{t-2} + \cdots + \beta^{t} v_0 vt​=gt​+βgt−1​+β2gt−2​+⋯+βtv0​

代入 v0=0v_0=0v0​=0:

vt=∑i=0t−1βi gt−i\boxed{v_t = \sum_{i=0}^{t-1} \beta^{i}\, g_{t-i}} vt​=i=0∑t−1​βigt−i​​

这个展开式是理解 Momentum 的关键:vtv_tvt​ 是过去所有梯度的指数加权平均。

权重分布(β=0.9\beta=0.9β=0.9 时):

梯度权重
gtg_tgt​(最新)1
gt−1g_{t-1}gt−1​0.9
gt−2g_{t-2}gt−2​0.81
...↓\downarrow↓
gt−t′g_{t-t'}gt−t′​0.9t′0.9^{t'}0.9t′

有效窗口长度约 11−β=10\frac{1}{1-\beta} = 101−β1​=10(β=0.9\beta=0.9β=0.9 时)。

2.2 为什么能抑制震荡

想象梯度序列 +1,−1,+1,−1,…+1, -1, +1, -1, \dots+1,−1,+1,−1,…:

  • 无 Momentum:净位移 =(+1−1)×η=0= (+1-1)\times \eta = 0=(+1−1)×η=0,原地打转
  • 有 Momentum(β=0.9\beta=0.9β=0.9):震荡被加权平均掉,vvv 收敛到一个非零的小正值,持续朝一个方向前进

本质:Momentum 是一个低通滤波器,把高频震荡滤掉,保留低频的漂移方向。

这就是为什么在陡峭峡谷(梯度方向来回变化)里,Momentum 能显著加速。

2.3 Nesterov Accelerated Momentum

改进:不用「当前」梯度算,而是用「位置更新后的」梯度。

θt=θt−1−η(βvt−1+gt(θt−ηvt−1))\theta_t = \theta_{t-1} - \eta \big(\beta v_{t-1} + g_t(\theta_t - \eta v_{t-1})\big) θ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−ηrtgtr_t = r_{t-1} + g_t^2, \qquad \theta_t = \theta_{t-1} - \frac{\eta}{\sqrt{r_t}} g_t rt​=rt−1​+gt2​,θt​=θt−1​−rt​​η​gt​

问题:rtr_trt​ 只增不减,分母越来越大 → lr 趋于 0,过早停止学习。

3.2 RMSProp:改成滑动平均

st=βst−1+(1−β)gt2,θt=θt−1−ηst+ϵgts_t = \beta s_{t-1} + (1-\beta) g_t^2, \qquad \theta_t = \theta_{t-1} - \frac{\eta}{\sqrt{s_t} + \epsilon} g_t st​=βst−1​+(1−β)gt2​,θt​=θt−1​−st​​+ϵη​gt​

两个改动:

  1. gt2g_t^2gt2​ 改成滑动平均 → sts_tst​ 可以增减,不再单调增
  2. 分母用 st\sqrt{s_t}st​​(标准差的估计)而非 rt\sqrt{r_t}rt​​(平方和)

直觉:参数变动剧烈(梯度大)→ 除以大数 → 步长自动变小。每个参数获得自己的学习率。

这就是「自适应学习率」的核心思想。


四、Adam:动量 + 自适应学习率

4.1 公式

mt=β1mt−1+(1−β1)gt(一阶矩:动量)vt=β2vt−1+(1−β2)gt2(二阶矩:方差)m^t=mt1−β1t(偏差修正)v^t=vt1−β2tθt=θt−1−ηm^tv^t+ϵ\begin{aligned}m_t &= \beta_1 m_{t-1} + (1-\beta_1) g_t \quad \text{(一阶矩:动量)}\\ v_t &= \beta_2 v_{t-1} + (1-\beta_2) g_t^2 \quad \text{(二阶矩:方差)}\\ \hat{m}_t &= \frac{m_t}{1-\beta_1^t} \quad \text{(偏差修正)}\\ \hat{v}_t &= \frac{v_t}{1-\beta_2^t} \\ \theta_t &= \theta_{t-1} - \eta \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} \end{aligned} mt​vt​m^t​v^t​θt​​=β1​mt−1​+(1−β1​)gt​(一阶矩:动量)=β2​vt−1​+(1−β2​)gt2​(二阶矩:方差)=1−β1t​mt​​(偏差修正)=1−β2t​vt​​=θt−1​−ηv^t​​+ϵm^t​​​

4.2 两个状态的直觉

状态统计什么类比
mtm_tmt​梯度的均值(方向)「最近梯度大致指向哪」
vtv_tvt​梯度的平方均值(大小)「梯度波动多大」

更新规则 = 沿梯度的平均方向走,步长按波动幅度缩放。

更新量=η⋅方向尺度\text{更新量} = \eta \cdot \frac{\text{方向}}{\text{尺度}} 更新量=η⋅尺度方向​

如果某个参数的梯度一直是 5.0(方向一致、尺度大),更新量 ≈η⋅1\approx \eta \cdot 1≈η⋅1,步长正常。 如果梯度是 +5,−5+5, -5+5,−5 交替(方向不一致),m≈0m \approx 0m≈0 → 更新量趋于 0,自动减小步长。

这就是 Adam 比 SGD 更快的原因:它自动识别并抑制震荡方向的更新。

4.3 偏差修正为什么必需(重要)

m0=v0=0m_0 = v_0 = 0m0​=v0​=0,所以第一步:

m1=(1−β1)g1m_1 = (1-\beta_1) g_1 m1​=(1−β1​)g1​

这明显低估了梯度的真实均值(应该是 g1g_1g1​,但只得到了 0.1g10.1g_10.1g1​)。早期所有估计都偏向 0。

修正:除以 1−βt1-\beta^t1−βt(因为 ∑i=0t−1(1−β)βi=1−βt\sum_{i=0}^{t-1}(1-\beta)\beta^i = 1-\beta^t∑i=0t−1​(1−β)βi=1−βt,归一化因子):

步数 ttt1−β1t1-\beta_1^t1−β1t​(β1=0.9\beta_1=0.9β1​=0.9)修正倍数 1/(1−βt)1/(1-\beta^t)1/(1−βt)
10.110×
20.195.3×
100.651.5×
1000.99997≈1.0

所以 Adam 的前几步更新量被放大了最多 10 倍。 这个修正保证了 lr 在任何时刻都是「名义 lr」,不会被初始化效应干扰。

这也解释了 warmup 的一个理由:Adam 早期更新量本来就偏大(即使修正后),warmup 再进一步慢慢提 lr,能让早期训练更稳。

4.4 Adam 的致命缺陷:权重衰减不兼容

关键问题:weight decay 在 SGD 里是加到梯度上的:

θ←θ−η(∇L+λθ)=(θ−ηλθ)−η∇L\theta \leftarrow \theta - \eta(\nabla L + \lambda\theta) = (\theta - \eta\lambda\theta) - \eta\nabla L θ←θ−η(∇L+λθ)=(θ−ηλθ)−η∇L

这个顺序很重要——先衰减,再走梯度。

但 Adam 里两者混在一起了:

θ←θ−η⋅m^tv^t+ϵ\theta \leftarrow \theta - \eta \cdot \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} θ←θ−η⋅v^t​​+ϵm^t​​

如果 weight decay 加进梯度 gtg_tgt​,它也会进入 mtm_tmt​ 和 vtv_tvt​:

mt=β1mt−1+(1−β1)(∇L+λθ)m_t = \beta_1 m_{t-1} + (1-\beta_1)(\nabla L + \lambda\theta) mt​=β1​mt−1​+(1−β1​)(∇L+λθ)

后果:

  1. 衰减量被 1v^t\frac{1}{\sqrt{\hat{v}_t}}v^t​​1​ 缩放。如果某个参数的梯度一直很小(v^t\hat{v}_tv^t​ 很小),衰减量会被放大——本来想轻微收缩,结果剧烈收缩
  2. 衰减量与梯度历史耦合,无法控制

这就是 AdamW(Adam with decoupled weight decay)的动机:把衰减从梯度里拿出来,直接作用于参数:

θt=θt−1−ηm^tv^t+ϵ−ηλθt−1\theta_t = \theta_{t-1} - \eta \frac{\hat{m}_t}{\sqrt{\hat{v}_t} + \epsilon} - \eta\lambda\theta_{t-1} θt​=θt−1​−ηv^t​​+ϵm^t​​−ηλθt−1​

衰减项不再进入自适应缩放,永远是恒定的 ηλθ\eta\lambda\thetaηλθ。

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 系更宽容,不用精调 lr
1
2
3
4
5
6
7
8
9
10
11
12
13
14

为什么 Transformer 用 AdamW 而不是 SGD?

  1. 各种尺度的注意力参数需要不同 lr,SGD 处理不好
  2. 稀疏梯度(部分维度梯度为0)下Adam 更稳
  3. Transformer 训练对小 lr 很敏感,Adam 的自适应让调参更容易

为什么视觉任务仍偏爱 SGD?

一个被广泛引用的经验:AdamW 类方法在训练后期会明显过拟合验证集,而 SGD+momentum 泛化更好。可能的解释是自适应方法的有效步长不稳定。文献很多、结论不完全统一,实践中以验证集结果为准。


六、β2 的一个实用细节

默认值 β2=0.999\beta_2 = 0.999β2​=0.999 在 LLM 训练中通常改成 0.95。 为什么?

11−β2\frac{1}{1-\beta_2}1−β2​1​ = 有效窗口长度:

β2\beta_2β2​窗口长度
0.9991000 步
0.9520 步
0.99100 步

LLM 训练中 β2=0.999\beta_2 = 0.999β2​=0.999 的窗口太长:因为训练过程中梯度分布本身在快速变化(从初期的大幅下降到后期的精细),用 1000 步的历史平均会让 vtv_tvt​ 严重滞后于当前状态。

改成 0.95(20 步窗口)后,二阶矩能更快跟上梯度的当前尺度。这是 HuggingFace、LLaMA 训练脚本的默认配置,也是现在的事实标准。


七、动手实验

实验 1:亲手实现 SGD / Momentum / Adam

比调torch.optim 有价值得多——你会真正理解每个优化器的「状态」是什么。

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

实测输出:

      SGD: 起点 1.0 → 最终 +1.0000     ← ★ 净位移为 0,完全原地打转
 Momentum: 起点 1.0 → 最终 +0.6915     ← 动量把它推过去了
     Adam: 起点 1.0 → 最终 +0.8455     ← 自适应缩放,步子更小但也有效
1
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 的偏差修正

python
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(正确)")
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   ← ★ 精确等于 lr
1
2

注意:不做修正时步长是 0.3162,比正确值 0.1 大了 3 倍,不是「缩小」。

为什么反直觉? 因为 ϵ\epsilonϵ 也在起作用。看清楚 β2=0.999\beta_2 = 0.999β2​=0.999 时发生了什么:

修正项值
m1m_1m1​0.1×3=0.30.1 \times 3 = 0.30.1×3=0.3
v1v_1v1​(未修正)0.001×9=0.0090.001 \times 9 = 0.0090.001×9=0.009
v1v_1v1​(修正后)0.009/(1−0.999)=9.00.009 / (1-0.999) = 9.00.009/(1−0.999)=9.0

v1未修正=0.0949,v1修正=3.0\sqrt{v_1^{\text{未修正}}} = 0.0949, \qquad \sqrt{v_1^{\text{修正}}} = 3.0 v1未修正​​=0.0949,v1修正​​=3.0

分子 m^\hat{m}m^ 从 0.3 → 3.0(放大 10 倍),分母也从 0.095 → 3.0(放大约 31 倍)。分母放得更多,所以最终步长偏小——0.1 × 0.3/0.095 = 0.316,正是实测值。

一句话总结偏差修正的作用:让第一步的有效步长精确等于名义 lr,不被初始化时的 β\betaβ 衰减拖慢。

实验 3:Adam vs AdamW 的差异(理解解耦的价值)

python
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}")
1
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.676259
1
2

梯度恒为 0.01(很小),weight decay=0.1,目标是把参数拉向 0:

  • AdamW:衰减量精确是 ηλp=0.1×0.1×p\eta\lambda p = 0.1 \times 0.1 \times pηλp=0.1×0.1×p,行为完全可预测
  • Adam:衰减量被送进 vtv_tvt​,最终的步长 ∝1v^t\propto \frac{1}{\sqrt{\hat v_t}}∝v^t​​1​ 包含了衰减的历史,于是衰减强度和梯度历史纠缠在一起

数值上这里只差 3%,看起来不多。但在真实训练里差异会累积:Adam 的 vtv_tvt​ 被污染后,会影响所有参数的自适应步长,而不只是被衰减的那些——这才是问题的严重性所在。

想看更戏剧性的差异,把梯度再调小试试(v^t\hat v_tv^t​ 越小,1/v^t1/\sqrt{\hat v_t}1/v^t​​ 的放大效应越强)。

实验 4:β2 的影响

python
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("原因:训练中梯度尺度快速变化,需要二阶矩跟上当前状态")
1
2
3
4
5
6

八、自测题

Q1:Adam 的 ϵ\epsilonϵ 是干什么的?如果设成 0 会有什么问题?

答案

防止除零。当某个参数在整个训练过程中梯度都是 0(可能因为它被 mask 了、或者初始化为 0 后从未更新),vt=0v_t = 0vt​=0,则 mtvt\frac{m_t}{\sqrt{v_t}}vt​​mt​​ 是 0/0。

加了 ϵ\epsilonϵ(默认 1e-8)后变成 mtvt+ϵ\frac{m_t}{\sqrt{v_t}+\epsilon}vt​​+ϵmt​​,分母有下界。

设成 0 的风险:

  1. 分母为 0 → NaN
  2. 即使不为 0,ϵ\epsilonϵ 太小会放大数值误差
  3. 实践中 ϵ\epsilonϵ 的选择对结果影响很小,因为它加在 vt\sqrt{v_t}vt​​ 上(vtv_tvt​ 通常远大于 1e-16)

Q2:为什么 SGD 在图像任务上泛化性常优于 Adam?

答案

没有公认定论,文献很多,主要假说:

  1. 自适应方法的等效学习率不稳定。Adam 的 η/vt\eta/\sqrt{v_t}η/vt​​ 在梯度尺度变化时会剧烈波动,导致它找到的是「容易优化的解」而非「泛化好的解」(类似 sharpness-aware minimization 的视角)。
  2. Adam 的隐式正则。m/vm/\sqrt{v}m/v​ 相当于对梯度做归一化,抹平了不同参数方向的尺度差异,导致模型倾向于找到「平坦但可能不泛化」的方向。
  3. 训练后期行为差异。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.001m \to 0.001m→0.001,v→0.001\sqrt{v} \to 0.001v​→0.001,更新量 =0.1×0.0010.001=0.1= 0.1 \times \frac{0.001}{0.001} = 0.1=0.1×0.0010.001​=0.1
  • 参数B(梯度一直 100):m→100m \to 100m→100,v→100\sqrt{v} \to 100v​→100,更新量 =0.1×100100=0.1= 0.1 \times \frac{100}{100} = 0.1=0.1×100100​=0.1

两个参数走一样多! 这就是自适应学习率的意义——不管梯度绝对大小如何,每个参数每步都走大致相同的距离。

这个特性对 Transformer 特别重要:注意力层的 Q/K/V 矩阵、FFN 的两层,梯度尺度差异巨大(可能相差几个数量级)。SGD 处理不好,Adam 能自动处理。

Q4:为什么 LLM 都用 warmup?如果没有 warmup 会怎样?

答案

warmup 的作用:在训练最初的几百/几千步把 lr 从 0 线性提升到目标值。

没有 warmup 的问题:

  1. 初期梯度不稳。训练开始时参数还没进入「有意义的空间」,梯度方向噪声大。直接用大 lr 会把参数推到糟糕的区域,一旦模型已经学到的结构被破坏,后面很难恢复。

  2. 自适应优化的二阶矩估计不准。mt,vtm_t, v_tmt​,vt​ 在头几步是严重有偏的(虽然有偏差修正,但修正本身让前几步的有效 lr 偏大),此时用满 lr 不稳。

  3. 配合 AdamW 会更糟。因为 weight decay 的实际强度是 ηλθ\eta\lambda\thetaηλθ——如果 lr 一开始就是满值,衰减也是满的,参数会在最初几步被剧烈拉向 0。

warmup 之后:

  • 梯度分布稳定了
  • 二阶矩估计准确了
  • 再配合 cosine decay,让 lr 在训练后期精细收敛

LLaMA 的做法:warmup 2000 步 + cosine decay到 10% 峰值。这套配置几乎成了标准。

Q5:PyTorch 里 Adam(lr=..., weight_decay=...) 和 AdamW 到底差在哪?如果我用了 Adam + weight_decay,我的模型会有什么问题?

答案

差在weight decay 的施加位置:

python
# 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            # 衰减独立施加
1
2
3
4
5
6
7
8
9
10
11

实际问题:

  1. 衰减项被 1v^\frac{1}{\sqrt{\hat v}}v^​1​ 缩放。梯度历史小的参数,衰减被放大——本意轻微收缩,实际剧烈收缩
  2. 衰减量和梯度历史纠缠,无法单独控制强度
  3. 在 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                     +b2
1
2

前向:

z1=W1x+b1a1=ReLU(z1)=max⁡(0,z1)z2=W2a1+b2L=(z2−y)2\begin{aligned} z_1 &= W_1 x + b_1 \\ a_1 &= \text{ReLU}(z_1) = \max(0, z_1) \\ z_2 &= W_2 a_1 + b_2 \\ L &= (z_2 - y)^2 \end{aligned} z1​a1​z2​L​=W1​x+b1​=ReLU(z1​)=max(0,z1​)=W2​a1​+b2​=(z2​−y)2​

现在要算 ∂L∂W1\frac{\partial L}{\partial W_1}∂W1​∂L​、∂L∂W2\frac{\partial L}{\partial W_2}∂W2​∂L​、∂L∂b1\frac{\partial L}{\partial b_1}∂b1​∂L​、∂L∂b2\frac{\partial L}{\partial b_2}∂b2​∂L​。

这里的关键认知:反向传播不是「每个参数独立算一次」,而是从输出往输入传一个「敏感度」,沿途用乘法法则分发给每个分支。


二、反向传播:一个参数一个参数推

2.1 从损失开始

∂L∂z2=2(z2−y)\frac{\partial L}{\partial z_2} = 2(z_2 - y) ∂z2​∂L​=2(z2​−y)

这个 2 就是 MSE 的特征。如果 loss 写成 12(z2−y)2\frac{1}{2}(z_2-y)^221​(z2​−y)2(很多教材这么写),梯度就干净了:∂L∂z2=(z2−y)\frac{\partial L}{\partial z_2} = (z_2-y)∂z2​∂L​=(z2​−y)。这就是为什么有些教材的 loss 系数是 12N\frac{1}{2N}2N1​ 而不是 1N\frac{1}{N}N1​ ——纯粹为了消掉这个 2。

2.2 分发给 W2W_2W2​ 和 b2b_2b2​

z2=W2a1+b2z_2 = W_2 a_1 + b_2z2​=W2​a1​+b2​,求 ∂L∂W2\frac{\partial L}{\partial W_2}∂W2​∂L​:

z2=∑j(W2)ij(a1)j+(b2)iz_2 = \sum_j (W_2)_{ij}(a_1)_j + (b_2)_i z2​=j∑​(W2​)ij​(a1​)j​+(b2​)i​

对 (W2)ij(W_2)_{ij}(W2​)ij​ 求导,(W2)ij(W_2)_{ij}(W2​)ij​ 只出现在第 i 项里:

∂L∂(W2)ij=∂L∂z2⋅∂z2∂(W2)ij=∂L∂z2⋅(a1)j\frac{\partial L}{\partial (W_2)_{ij}} = \frac{\partial L}{\partial z_2}\cdot\frac{\partial z_2}{\partial (W_2)_{ij}} = \frac{\partial L}{\partial z_2}\cdot (a_1)_j ∂(W2​)ij​∂L​=∂z2​∂L​⋅∂(W2​)ij​∂z2​​=∂z2​∂L​⋅(a1​)j​

写成矩阵形式(避免下标地狱):

∂L∂W2=∂L∂z2⏟[1,1]⋅a1⊤⏟[1,H]=outer product\boxed{\frac{\partial L}{\partial W_2} = \underbrace{\frac{\partial L}{\partial z_2}}_{\text{[1,1]}} \cdot \underbrace{a_1^\top}_{\text{[1,H]}} = \text{outer product}} ∂W2​∂L​=[1,1]∂z2​∂L​​​⋅[1,H]a1⊤​​​=outer product​

外积的形状:[out,1]×[1,in]=[out,in][\text{out}, 1] \times [1, \text{in}] = [\text{out}, \text{in}][out,1]×[1,in]=[out,in] ✓ 正好是 W2W_2W2​ 的形状。

∂L∂b2=∂L∂z2\frac{\partial L}{\partial b_2} = \frac{\partial L}{\partial z_2} ∂b2​∂L​=∂z2​∂L​

(因为 b2b_2b2​ 是加法,导数为 1)

2.3 穿过 ReLU

∂a1∂z1={1z1>00z1≤0\frac{\partial a_1}{\partial z_1} = \begin{cases} 1 & z_1 > 0 \\ 0 & z_1 \le 0 \end{cases} ∂z1​∂a1​​={10​z1​>0z1​≤0​

所以:

∂L∂z1=∂L∂a1⊙1[z1>0]\frac{\partial L}{\partial z_1} = \frac{\partial L}{\partial a_1} \odot \text{1}[z_1 > 0] ∂z1​∂L​=∂a1​∂L​⊙1[z1​>0]

这里有个重要事实:ReLU 会把负半轴的梯度完全掐断(置0)。 这是 ReLU 「死神经元」问题的根源——如果一个神经元对所有输入都是负的,它的梯度永远是 0,参数永远不更新,等于废掉了。

2.4 分发给 W1W_1W1​ 和 b1b_1b1​

∂L∂W1=∂L∂z1⋅x⊤,∂L∂b1=∂L∂z1\frac{\partial L}{\partial W_1} = \frac{\partial L}{\partial z_1} \cdot x^\top, \qquad \frac{\partial L}{\partial b_1} = \frac{\partial L}{\partial z_1} ∂W1​∂L​=∂z1​∂L​⋅x⊤,∂b1​∂L​=∂z1​∂L​

到这里就完成了。 完整链条:

∂L∂z2→∂L∂a1→∂L∂z1→∂L∂W1,∂L∂b1\frac{\partial L}{\partial z_2} \to \frac{\partial L}{\partial a_1} \to \frac{\partial L}{\partial z_1} \to \frac{\partial L}{\partial W_1}, \frac{\partial L}{\partial b_1} ∂z2​∂L​→∂a1​∂L​→∂z1​∂L​→∂W1​∂L​,∂b1​∂L​


三、通用规则(记住这三条就够)

反向传播就是从后往前反复应用这三条规则:

规则 1:加法 → 梯度直接分发

z = a + b
∂L/∂a = ∂L/∂z    ∂L/∂b = ∂L/∂z
1
2

多个分支收到相同的上游梯度(所以 GradienT 累加不是 bug,是数学要求)。

规则 2:乘法 → 梯度交叉相乘

z = a × b
∂L/∂a = ∂L/∂z × b    ∂L/∂b = ∂L/∂z × a
1
2

规则 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]
1
2
3
4
5

记忆方法:「哪个输入和梯度同形状,就用它的转置去乘另一个」。

  • WWW 是 [in, out],要得到 [in, out],就用 a⊤[B,in]a^\top[B,in]a⊤[B,in] 乘 (∂L/∂z)[B,out](\partial L/\partial z)[B,out](∂L/∂z)[B,out]
  • aaa 是 [B, in],要得到 [B, in],就用 (∂L/∂z)[B,out](\partial L/\partial z)[B,out](∂L/∂z)[B,out] 乘 W⊤[out,in]W^\top[out,in]W⊤[out,in]

这就是外积的批量版:∂L/∂W=∑i(第i 个样本的 ∂L/∂zi)⊗ai\partial L/\partial W = \sum_i (\text{第}i\text{ 个样本的}\ \partial L/\partial z_i) \otimes a_i∂L/∂W=∑i​(第i 个样本的 ∂L/∂zi​)⊗ai​ —— 每个样本贡献一个外积,加起来就是矩阵乘。

PyTorch 里的对应关系:

数学写法PyTorch
外积 (ab⊤)(ab^\top)(ab⊤)torch.outer(a, b) 或 a @ b.T
张量对元素积 ⊙\odot⊙a * b
矩阵乘@
沿 dim 求和.sum(dim=n)

一句话总结

反向传播 = 反向应用链式法则 + 把「敏感度」沿计算图分发。 加法节点等值分发,乘法节点交叉相乘,矩阵乘用外积。


四、梯度消失与爆炸的数学本质

这是本文最重要的部分。理解了这段,你就能自己判断任何架构设计的梯度行为,不需要背「ResNet 解决了梯度消失」这种结论。

4.1 梯度经过一层的衰减/放大

考虑一个 LLL 层、宽度 nnn 的全连接网络(ReLU + 合适的初始化)。每层权重的 Jacobian 矩阵 JiJ_iJi​ 的元素典型大小约 O(1n)O(\frac{1}{\sqrt n})O(n​1​)(He 初始化的设计目标)。

反向传播时梯度要连乘 LLL 个 Jacobian:

∂L∂x=JL⋅JL−1⋯J1⋅∂L∂out\frac{\partial L}{\partial x} = J_L \cdot J_{L-1} \cdots J_1 \cdot \frac{\partial L}{\partial \text{out}} ∂x∂L​=JL​⋅JL−1​⋯J1​⋅∂out∂L​

每个因子贡献一个 1n\frac{1}{\sqrt n}n​1​,连乘 LLL 次:

(1n)L=n−L/2\left(\frac{1}{\sqrt n}\right)^L = n^{-L/2} (n​1​)L=n−L/2

网络LLLnnnn−L/2n^{-L/2}n−L/2
浅层310010−310^{-3}10−3
深层3010010−1510^{-15}10−15
深层100100010−15010^{-150}10−150

10−1510^{-15}10−15 在 float32 的精度下就是 0。 这就是梯度消失。

4.2 Sigmoid 的额外问题

如果激活函数是 sigmoid/tanh,导数最大只有 0.25(sigmoid):

σ′(z)=σ(z)(1−σ(z))≤0.25\sigma'(z) = \sigma(z)(1-\sigma(z)) \le 0.25 σ′(z)=σ(z)(1−σ(z))≤0.25

每经过一层,梯度至少乘以 0.25。 10 层就是 0.2510≈10−60.25^{10} \approx 10^{-6}0.2510≈10−6。

为什么 ReLU 缓解了这个问题:ReLU′(z)≥0\text{ReLU}'(z) \ge 0ReLU′(z)≥0 且期望为 1(对 z>0z>0z>0 部分导数为 1),所以不会引入额外的衰减。这就是 2012 年 ReLU 带来突破的数学原因。

4.3 梯度爆炸

反过来,如果权重初始化太大(JJJ 的元素 ≫1\gg 1≫1),连乘后会指数爆炸:

∂L∂x∝(something≫1n)L→∞\frac{\partial L}{\partial x} \propto \left(\frac{\text{something} \gg 1}{\sqrt n}\right)^L \to \infty ∂x∂L​∝(n​something≫1​)L→∞

表现:loss 突然变成 nan、梯度值离谱、参数被更新到溢出。

解法是梯度裁剪(gradient clipping):

python
torch.nn.utils.clip_grad_norm_(model.parameters(), max_norm=1.0)
1

把所有梯度作为一个整体,计算 L2 范数,如果超过阈值就等比缩放:

g^=g⋅max_norm∥g∥当∥g∥>max_norm\hat{g} = g \cdot \frac{\text{max\_norm}}{\|g\|} \quad \text{当} \|g\| > \text{max\_norm} g^​=g⋅∥g∥max_norm​当∥g∥>max_norm

注意是等比缩放所有梯度,不是逐元素裁剪(后者会改变梯度方向,破坏优化语义)。

LLM 训练标准配置里都有 clip=1.0。这不是可选项。

4.4 残差连接如何解决

残差连接:al+1=al+f(al)a_{l+1} = a_l + f(a_l)al+1​=al​+f(al​)

反向传播时,梯度有一条「直达通路」:

∂L∂al=∂L∂al+1(I+∂f∂al)\frac{\partial L}{\partial a_l} = \frac{\partial L}{\partial a_{l+1}}\left(I + \frac{\partial f}{\partial a_l}\right) ∂al​∂L​=∂al+1​∂L​(I+∂al​∂f​)

那个 III(单位矩阵)意味着梯度可以原封不动地传下去。

即使 fff 那部分的梯度衰减到 0,还有 III 兜着。所以:

∂L∂a1=∏l=1L(I+Jl)\frac{\partial L}{\partial a_1} = \prod_{l=1}^{L}\left(I + J_l\right) ∂a1​∂L​=l=1∏L​(I+Jl​)

这个乘积里,每项都含 III,展开后至少有一项全是 III 的乘积(对应「所有残差路径都不走 fff」那条路线),所以梯度不衰减。

这是 ResNet 能训 100 层的数学保证。 详见第 8 篇。

4.5 一个必须掌握的自测方法

判断一个新架构会不会梯度消失/爆炸,不要靠猜——做实验测量。

python
# 逐层测量梯度范数,一眼看出是否衰减/爆炸
for name, param in model.named_parameters():
    if param.grad is not None:
        print(f"{name:40s} ||grad|| = {param.grad.norm().item():.3e}")
1
2
3
4

判读:

  • 逐层指数下降(如 1e-1 → 1e-3 → 1e-5 → 1e-8)→ 梯度消失
  • 逐层指数上升 → 梯度爆炸
  • 各层量级相当(1e-2 量级上下浮动)→ 健康

这个技巧在任何论文的复现里都用得上。看到一个新架构,第一件事就是打这个表。


五、动手实验

实验 1:验证手推公式和 autograd 完全一致

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

实测输出:

形状对齐(必须和参数完全一致):
  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+00
1
2
3
4
5
6
7
8
9
10
11

误差全是 0(浮点运算顺序一致时)。这证明了你的推导和 autograd 做的是同一件事。

⚠️ 手推时我踩的三个坑(重要)

这个实验的价值一半在于踩坑。三个错误都是我第一次写时犯的:

坑 1:忘了 batch 平均的系数

loss = (z2-y)**2.mean() 是对 batch 和输出维都求平均,所以:

∂L∂z2=2(z2−y)B\frac{\partial L}{\partial z_2} = \frac{2(z_2 - y)}{B} ∂z2​∂L​=B2(z2​−y)​

我一开始写成 2*(z2-y),漏了除以 B=4,结果所有梯度差 4 倍。

坑 2:矩阵乘法顺序写反

z1=x@W1z_1 = x @ W_1z1​=x@W1​ 中 W1W_1W1​ 是 [3,2](in×out),所以:

∂L∂W1=x⊤⋅∂L∂z1(先转 x)\frac{\partial L}{\partial W_1} = x^\top \cdot \frac{\partial L}{\partial z_1} \quad (\text{先转 } x) ∂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:测量梯度消失

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

你会观察到:ReLU 网络的梯度衰减远慢于 Sigmoid;层数增加时两边都衰减,但 Sigmoid 更严重。这就是 ReLU 优于 Sigmoid 的定量证据。

实验 3:残差连接如何消除衰减

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

你会看到:深度越大,普通网络的梯度衰减倍数呈指数增长,而残差网络几乎保持恒定。这就是 ResNet 能训 100+ 层的定量证据。


六、自测题

Q1:为什么 ReLU 会出现「死神经元」?什么条件下会出现?

答案

机制:ReLU 的梯度是 1[z>0]\text{1}[z>0]1[z>0]。如果某个神经元对所有输入都输出 z≤0z \le 0z≤0,它的梯度永远是 0,参数永远不更新。

什么条件会触发:

  1. 初始化不当:权重初始化为全 0 或负偏置过大,初始输出就大面积为负
  2. 学习率过大:一次更新把参数推到负半轴,之后再也回不来
  3. 输入分布本身为负(比如前面接了负偏置、或者数据本身偏负)

为什么回不来:一旦某个神经元对所有输入都是负输出,它就不会参与任何有用的计算,其他神经元也不会因为它而得到有用的梯度(梯度是累加的,但这个神经元贡献 0),所以整个系统失去了这个神经元。这就是「死」。

解法:

  • Leaky ReLU:max⁡(0.05z,z)\max(0.05z, z)max(0.05z,z),给负半轴一个小的非零斜率
  • 参数化 ReLU(PReLU):斜率可学习
  • 正确初始化(He 初始化,让初始输出正负各半)
  • 合适的 lr

Q2:如果网络所有权重都初始化为 0,会发生什么?

答案

输出恒为 0,且梯度也为 0,训练完全不进行。

推导:

  • 前向:z=Wx+bz = Wx + bz=Wx+b,若 W=0W=0W=0 则 z=0z=0z=0(不管 bbb 多大),所有层输出 0
  • 反向:∂L∂W2=∂L∂z2⋅a1⊤\frac{\partial L}{\partial W_2} = \frac{\partial L}{\partial z_2} \cdot a_1^\top∂W2​∂L​=∂z2​∂L​⋅a1⊤​,而 a1=0a_1 = 0a1​=0,所以 ∂L∂W2=0\frac{\partial L}{\partial W_2} = 0∂W2​∂L​=0
  • 逐层回推,所有梯度都是 0

对称性问题:如果所有神经元用同一个初始化且输入也一样,它们的梯度完全相同 → 参数更新也相同 → 永远保持对称,等效于一个神经元。

注意区分:

  • 权重 W 不能初始化为 0
  • 偏置 b 可以初始化为 0(因为它不造成对称性问题)
  • 最后一个全连接层的 bias 可以初始化为 0,但它的 weight 初始化为 0 会导致 logits 全相同 → 初始 loss 就是 log⁡K\log KlogK

这也是为什么 PyTorch 里的 nn.Linear 默认 bias=True 且初始化均匀。

Q3:梯度裁剪是把超过阈值的部分「截断」,为什么等比缩放(全梯度乘一个系数)更好?

答案

逐元素裁剪(clip by value):

gi←max⁡(−ϵ,min⁡(ϵ,gi))g_i \leftarrow \max(-\epsilon, \min(\epsilon, g_i)) gi​←max(−ϵ,min(ϵ,gi​))

这会改变梯度的方向。比如原梯度是 (0.1,10.0)(0.1, 10.0)(0.1,10.0),裁剪后是 (0.1,1.0)(0.1, 1.0)(0.1,1.0),方向从近似 y 轴变成更接近 x 轴。优化的方向被改掉了——这违反了「梯度下降沿最陡下降方向」的前提。

等比缩放(clip by norm):

g^=g⋅ϵ∥g∥\hat{g} = g \cdot \frac{\epsilon}{\|g\|} g^​=g⋅∥g∥ϵ​

所有分量乘同一个系数,方向完全不变,只是整体步长被限制。

实现:

python
torch.nn.utils.clip_grad_norm_(
    model.parameters(),
    max_norm=1.0,
    error_if_nonfinite=False,   # nan 时跳过
)
1
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,逐层打印梯度。如果这个基线是健康的,说明问题在你的具体设计(初始化、归一化、结构);如果基线也不健康,那就是深度本身带来的,需要残差连接。

python
# 最小复现
model = nn.Sequential(*[nn.Sequential(nn.Linear(64,64), nn.ReLU()) for _ in range(50)])
1
2

下一篇 → 归一化:从 BatchNorm 到 RMSNorm

上一级: 目录 · 上一篇

6 · 归一化:BatchNorm → RMSNorm

核心问题:归一化到底在解决什么?为什么 BatchNorm 在 CNN 里统治、在 NLP 里被淘汰?RMSNorm 又省了什么?


一、归一化解决的两个问题

1.1 尺度问题(forward)

考虑一个深网络里的激活值。如果第 lll 层的输出尺度是第 l−1l-1l−1 层的 kkk 倍(k>1k>1k>1),那么到第 LLL 层尺度就是 kLk^LkL。

k=1.1k = 1.1k=1.1、L=50L = 50L=50 时,kL≈117k^L \approx 117kL≈117。指数增长,一层层放大,网络极易饱和(sigmoid 饱和区 / ReLU 全死)。

归一化把每层的输出强制拉到「零均值单位方差」,尺度被控制住了。

1.2 优化问题(backward)

梯度里含因子 W⊤W^\topW⊤。如果 WWW 的谱半径(最大奇异值)ρ(W)>1\rho(W) > 1ρ(W)>1,梯度连乘会指数爆炸;ρ(W)<1\rho(W) < 1ρ(W)<1 则指数消失。

归一化让每层的 Jacobian 更接近「谱半径为 1 的等距映射」,梯度既不爆炸也不消失。

归一化的核心价值:让每一层的梯度尺度可控,从而让深层网络可训练。


二、BatchNorm 详解

2.1 公式

对一个 mini-batch 的同一通道,统计均值和方差,归一化后再缩放平移:

x^i=xi−μBσB2+ϵ,yi=γx^i+β\hat{x}_i = \frac{x_i - \mu_B}{\sqrt{\sigma_B^2 + \epsilon}}, \qquad y_i = \gamma \hat{x}_i + \beta x^i​=σB2​+ϵ​xi​−μB​​,yi​=γx^i​+β

其中(对 batch 维度 mmm 求):

μB=1m∑ixi,σB2=1m∑i(xi−μB)2\mu_B = \frac{1}{m}\sum_i x_i, \qquad \sigma_B^2 = \frac{1}{m}\sum_i (x_i - \mu_B)^2 μB​=m1​i∑​xi​,σB2​=m1​i∑​(xi​−μB​)2

γ,β\gamma, \betaγ,β 是可学习的(shape = 通道数),作用是「让网络可以自己决定要不要归一化、以及归一化到什么分布」。如果 γ=β=0\gamma=\beta=0γ=β=0,BN 就退化成恒等映射——所以加 BN 的位置总会初始化成 γ=1,β=0\gamma=1, \beta=0γ=1,β=0。

2.2 关键设计:训练/推理行为不同

推理时没有 batch 可用,怎么办?

用训练时积累的滑动平均(running mean / running var):

python
# PyTorch 内部维护
self.register_buffer('running_mean', torch.zeros(num_features))
self.register_buffer('running_var', torch.ones(num_features))
1
2
3

训练时更新:

running_mean=(1−momentum)⋅running_mean+momentum⋅μB\text{running\_mean} = (1-\text{momentum}) \cdot \text{running\_mean} + \text{momentum} \cdot \mu_B 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 的关键差异

BatchNormLayerNorm
统计维度batch 维 + 空间维(对整个 batch 统计)特征维(每个样本独立算)
依赖 batch 大小是否
训练/推理行为不同是否
适合序列否(batch小 + 变长序列)是
消融论文ResNetTransformer / LLaMA

3.2 公式

对单个样本的特征维度统计:

μ=1d∑j=1dxj,σ2=1d∑j=1d(xj−μ)2\mu = \frac{1}{d}\sum_{j=1}^{d} x_j, \qquad \sigma^2 = \frac{1}{d}\sum_{j=1}^{d}(x_j-\mu)^2 μ=d1​j=1∑d​xj​,σ2=d1​j=1∑d​(xj​−μ)2

x^=x−μσ2+ϵ,y=γ⊙x^+β\hat{x} = \frac{x - \mu}{\sqrt{\sigma^2+\epsilon}}, \qquad y = \gamma \odot \hat{x} + \beta x^=σ2+ϵ​x−μ​,y=γ⊙x^+β

注意:μ,σ\mu, \sigmaμ,σ 是对单个样本的所有特征算的,和其他样本无关。

3.3 为什么 NLP 必须用 LayerNorm

三个决定性原因:

  1. batch 统计量不可靠。NLP 的 batch size 常常是 1(长序列 + 大词表),BN 根本没足够样本算方差。
  2. 变长序列会引入 padding 污染。同一个 batch 里不同长度序列的 padding 位置不同,如果对 batch 统计,padding 的 0 值会污染均值方差。
  3. 推理行为必须确定性。变长输入每次凑到的 batch 不同,BN 的 running 统计不准。

LayerNorm 完全避开了这三个问题——每个 token 独立归一化,跟 batch 里有什么完全无关。

3.4 一个反直觉的现象

Transformer 里,LayerNorm 放在残差连接的「后面」(Post-LN)还是「前面」(Pre-LN)?

Post-LN(原始 Transformer)Pre-LN(现代 LLM)
结构x+Sublayer(x)x + \text{Sublayer}(x)x+Sublayer(x) 后再 LNSublayer(x)+x\text{Sublayer}(x) + xSublayer(x)+x 里 LN 在前
深层训练需要 warmup稳定,warmup 可选
最终性能可能更好略差或持平
代表原始 TransformerGPT-2/LLaMA 等几乎全部现代 LLM

Pre-LN 的残差路径是干净的恒等映射:

Pre-LN: xl+1=xl+F(LN(xl))\text{Pre-LN: } x_{l+1} = x_l + F(\text{LN}(x_l)) Pre-LN: xl+1​=xl​+F(LN(xl​))

梯度反向时 ∂xl+1∂xl=I+∂F∂LN⋅1σ\frac{\partial x_{l+1}}{\partial x_l} = I + \frac{\partial F}{\partial \text{LN}} \cdot \frac{1}{\sigma}∂xl​∂xl+1​​=I+∂LN∂F​⋅σ1​,那个 III 保证了梯度直通。深层训练立刻稳定。

Post-LN 则是 xl+1=LN(xl+F(xl))x_{l+1} = \text{LN}(x_l + F(x_l))xl+1​=LN(xl​+F(xl​)),LN 夹在残差路径中间,梯度要穿过 LN 才能回到前面。

这就是为什么 LLaMA 用 RMSNorm + Pre-LN(第 12 篇会详细讲)。


四、RMSNorm:LLaMA 的选择

4.1 省掉了什么

RMSNorm(LLaMA / T5 提出)观察到一个事实:LayerNorm 里,去掉均值中心化,效果几乎不掉。

RMSNorm(x)=xRMS(x)⊙γ,RMS(x)=1d∑jxj2\text{RMSNorm}(x) = \frac{x}{\text{RMS}(x)} \odot \gamma, \qquad \text{RMS}(x) = \sqrt{\frac{1}{d}\sum_j x_j^2} RMSNorm(x)=RMS(x)x​⊙γ,RMS(x)=d1​j∑​xj2​​

对比:

LayerNormRMSNorm
求均值 μ\muμ需要不需要
求方差需要(先减均值再平方)只需平方和的均方根
减均值需要不需要
可学习参数γ,β\gamma, \betaγ,β只有 γ\gammaγ

省掉的操作:均值计算、减法、ddd 个 beta 参数。

4.2 为什么省掉还能work

直觉解释:归一化的主要作用是「控制尺度」,不是「控制中心」。

LayerNorm 的 x−μσ2+ϵ\frac{x-\mu}{\sqrt{\sigma^2+\epsilon}}σ2+ϵ​x−μ​ 可以近似理解为:

除以「偏差」+除以「标准差」\text{除以「偏差」} + \text{除以「标准差」} 除以「偏差」+除以「标准差」

RMSNorm 只保留了第二项(除以 RMS ≈ 标准差)。实验表明第一项对效果贡献很小。

更深层的解释:Transformer 里每个 token 的表示本身已经有明确的语义中心,减去 batch/样本均值带来的收益有限;而那一步的数值开销和显存占用是实打实的。

4.3 实际收益(LLM 训练是显存瓶颈)

LayerNormRMSNorm
每层参数(d=4096)2d=81922d = 81922d=8192d=4096d = 4096d=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 的标准做法,现在所有视觉网络的实现里都能看到这段逻辑:

python
# 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 前的权重
1
2
3
4
5
6
7

这个「训练误差更高但测试误差更好」的现象很反直觉:weight decay 减小了拟合能力(训练误差上升),却提升了泛化(测试误差下降)。过拟合的减少不一定要靠训练损失更低来实现。

5.2 推理时的 batch 依赖

BN 的输出依赖 batch 里的其他样本。这导致:

  • 同一个输入,在不同 batch 组成下输出不同
  • 线上推理时 batch 大小变化 → 结果变化 → 难以复现
  • 这也是为什么很多线上服务要求「固定 batch size」

LayerNorm/RMSNorm 没有这个问题——每个样本独立计算。


六、动手实验

实验 1:BatchNorm 的 train/eval 差异

python
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 无关")
1
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 次输出是否唯一: True
1
2
3
4

PyTorch 直接帮你拦住了 batch=1 这个坑——这就是自测题 Q1 说的失效场景,框架层面已经做了防御。

实验 2:LayerNorm 与 batch 完全无关

python
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")
1
2
3
4
5
6
7
8
9
10
11
12
13

LayerNorm 只对自己样本的特征维做统计,所以输出与 batch 组成完全无关。这是它在 NLP 里不可替代的原因。

实验 3:RMSNorm 省掉了「中心化」

python
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}")
1
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.905
1
2
3

这个对比很说明问题:

  • LayerNorm:均值被拉到 0,标准差拉到 1
  • RMSNorm:保留了输入的偏移(+0.426),但标准差照样被控制住(0.905 ≈ 1)

结论:归一化的核心作用是「控制尺度」,不是「控制中心」。 RMSNorm 砍掉的正是次要功能,所以能省。

实验 4:归一化如何缓解梯度消失

python
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} 倍")
1
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 倍
1
2
3

放大了近 2000 万倍。这不是修辞——20 层不带归一化的网络,输入端梯度已经衰减到 10−810^{-8}10−8,在 float32 下几乎等同于 0,参数完全学不动。加上 LayerNorm 后梯度回到 10010^{0}100 量级,完全可用。

归一化把每层 Jacobian 的谱半径锚定在 1 附近,阻止了连乘时的指数衰减。这正是深层网络能训起来的前提。


七、自测题

Q1:为什么 BatchNorm 在 batch size = 1 时行为异常?

答案

训练时:BN 在单个样本上算均值方差,μB=x1\mu_B = x_1μB​=x1​,σB2=0\sigma_B^2 = 0σB2​=0,于是:

x^1=x1−x10+ϵ=0\hat{x}_1 = \frac{x_1 - x_1}{\sqrt{0 + \epsilon}} = 0 x^1​=0+ϵ​x1​−x1​​=0

输出恒为 β\betaβ,完全丢失了输入的信息。同时反向传播时 ∂y∂x^=1σ2+ϵ=1ϵ\frac{\partial y}{\partial \hat{x}} = \frac{1}{\sqrt{\sigma^2+\epsilon}} = \frac{1}{\sqrt\epsilon}∂x^∂y​=σ2+ϵ​1​=ϵ​1​,如果 ϵ\epsilonϵ 也很小,梯度会爆炸。

推理时:用 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 自己的 γ,β\gamma, \betaγ,β(1 维参数)做 weight decay 没有意义,只是让模型平移边界。

ResNet 论文的实验:

  • 全部层衰减:wd=1e-3 最佳
  • BN 前不衰减:wd=1e-4 更好,训练误差更高但测试误差更低

结论:正则化导致训练误差升高是正常的,重要的是泛化。「训练误差更低」和「模型更好」不是一回事。

这个技巧现在所有视觉网络都默认开启。

Q3:RMSNorm 去掉了减均值,为什么效果不掉?更深的原因是什么?

答案

表层原因:归一化有两个作用——控制尺度(除以标准差)和控制中心(减均值)。实验表明控制尺度的贡献占绝大部分。

深层原因:

  1. LN 的减均值在数学上和「白化」相关,但只在特征之间有强相关时才有意义。如果特征之间近似正交,减均值只是平移,收益很小。

  2. Transformer 里 LLM 已经有 RMSNorm 控制了极端值,不需要 LN 再做一次中心化。

  3. BERT 里 LN 的中心化作用被 embedding 的 LayerNorm 部分抵消——多层 LN 叠加时,效果已经饱和。

实证支持:

  • T5 论文(Google):把 LN 换成 RMSNorm,效果几乎不变,速度更快
  • LLaMA 论文:明确采用 RMSNorm,理由是「经验上更稳定」

一个补充观察:RMSNorm 在深层网络里有时更稳定。原因之一是它避免了「减均值后平方」这个数值上不稳定的操作链(当均值远大于标准差时,x−μx-\mux−μ 会有灾难性抵消)。

Q4:什么情况下 BatchNorm 反而比 LayerNorm 好?

答案

只有视觉任务,且 batch 足够大的时候。 具体:

  1. CNN 图像分类:ResNet / EfficientNet / ConvNeXt 全都用 BN。图像的空间维度大(224×224),即使 batch=32 也相当于统计了 32×224232 \times 224^232×2242 个样本,统计量非常可靠。

  2. 特定场景:AdaFace、TransNorm 等混合方案会判断:

    • 如果特征维度小、batch 大 → BN
    • 如果特征维度大、batch 小 → LN
  3. BatchNorm 有额外优势:

    • 训练/推理时的正则化效应(dropout 式的噪声)被证明对视觉任务有帮助
    • 可以用「冻结 running 统计量」的方式做模型校准和量化
    • 某些部署场景下 BN 可以折叠进前面的卷积层(推理时等价于一个普通 conv),LayerNorm 不行

判断标准:

场景选择
图像分类(batch≥32)BatchNorm
目标检测/分割(batch小、高分辨率)BatchNorm 或 GroupNorm
Transformer / NLPLayerNorm
大语言模型RMSNorm(+ Pre-LN)
小 batch 在线学习LayerNorm(BN 的 running 统计不可靠)

下一篇 → 卷积与感受野

上一级: 目录 · 上一篇

7 · 卷积与感受野

核心问题:卷积为什么比全连接更适合图像?什么条件下卷积等价于全连接?感受野怎么算?


一、卷积的数学定义

1.1 从「滑动窗口的点积」理解

二维卷积在计算机视觉里的实际计算是:

out[i,j]=∑u,vinput[i+u,j+v]⋅K[u,v]+b\text{out}[i,j] = \sum_{u,v} \text{input}[i+u, j+v] \cdot K[u,v] + b out[i,j]=u,v∑​input[i+u,j+v]⋅K[u,v]+b

注意这不是数学上的卷积(数学上是先翻转再相关),但深度学习里习惯叫卷积。

两个关键特性:

  1. 权重共享(weight sharing):同一个卷积核 KKK 在所有位置复用
  2. 局部连接(local connectivity):每个输出只依赖一个局部区域

1.2 参数量对比(这是卷积最大的优势)

设输入 224×224×3224 \times 224 \times 3224×224×3,输出 224×224×64224 \times 224 \times 64224×224×64:

全连接层(从 150528 维映射到 50176 维):

150528×50176≈7.5×109参数150528 \times 50176 \approx 7.5 \times 10^9 \quad \text{参数} 150528×50176≈7.5×109参数

卷积层(3×3 卷积,64 个输出通道):

3×3×3×64+64=1792参数3 \times 3 \times 3 \times 64 + 64 = 1792 \quad \text{参数} 3×3×3×64+64=1792参数

差4 百万倍,而卷积的表达能力并不弱(因为它利用了图像的局部性先验)。

这个参数量的差距是卷积在视觉领域统治的根本原因,不是「效果更好」,而是「能训得动」。


二、卷积 vs 全连接:等价条件

2.1 什么情况下卷积 == 全连接

如果卷积核的空间尺寸等于输入的空间尺寸,那么卷积就退化成全连接。

例:5×55 \times 55×5 的输入,用 5×55 \times 55×5 的卷积核(stride=1, padding=0)→ 输出 1×11 \times 11×1。

此时每个输出需要看到全部输入,权重共享的约束还在,但已经没有空间局部性可言了。

2.2 更精确的等价条件

3×33 \times 33×3 卷积(padding=1)在 H×WH \times WH×W 输入上产生 H×WH \times WH×W 输出。对输出位置 (i,j)(i,j)(i,j):

out[i,j]=W:,:,0⋅Xi−1:i+1,j−1:j+1+b\text{out}[i,j] = W_{:, :, 0}\cdot X_{i-1:i+1, j-1:j+1} + b out[i,j]=W:,:,0​⋅Xi−1:i+1,j−1:j+1​+b

每个输出用的都是同一个 WWW,但看的是不同的输入块。所以它不是全连接——全连接每个输出位置应该有不同的权重。

但是:可以把「卷积」看成「一种结构化的全连接」。WWW 被约束成 11 个不同的矩阵,每个矩阵的 9 个权重被共享。约束 = 正则化。

这就是卷积的泛化能力来源:用「权重必须共享」这个先验,替代了「参数独立」的自由度。


三、感受野(Receptive Field)

3.1 定义

感受野 = 输出特征图上一个元素所「看到」的输入区域大小。

这是理解 CNN 结构的核心工具。

3.2 计算公式

逐层递推:

rl=rl−1+(kl−1)⋅∏i=1l−1si,jl=jl−1+(kl−1)∏i=1l−1sir_l = r_{l-1} + (k_l - 1) \cdot \prod_{i=1}^{l-1} s_i, \qquad j_l = j_{l-1} + (k_l - 1)\prod_{i=1}^{l-1}s_i rl​=rl−1​+(kl​−1)⋅i=1∏l−1​si​,jl​=jl−1​+(kl​−1)i=1∏l−1​si​

其中 rrr 是感受野大小,jjj 是跳跃间隔(jump),kkk 是核大小,sss 是 stride。

stride=1 的简化公式:

rl=1+∑i=1l(ki−1)r_l = 1 + \sum_{i=1}^{l}(k_i - 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 511

注意感受野是累加的:1+5×2=111 + 5 \times 2 = 111+5×2=11。

核心洞察:用 5 层 3×3 卷积(感受野 11)比 1 层 11×11 卷积(感受野 11)好得多,因为:

  1. 中间有 4 次非线性 → 表达力更强
  2. 参数量:5 层 3×3 = 5×9=455 \times 9 = 455×9=45 倍权重 vs 1 层 11×11 = 121 倍
  3. 中间可以做下采样(VGG 的设计哲学)

3.3 感受野的三种叠加方式

设计网络时必须明确「想要多大的感受野」,然后选结构:

方式做法感受野增长速度代表网络
堆叠卷积连续多个 3×3线性(1+2L1+2L1+2L)VGG
Pooling用池化降采样平方级增长VGG / AlexNet
空洞卷积dilation > 1指数增长DeepLab / WaveNet

空洞卷积(Dilated / Atrous Convolution) 值得单独说:

感受野=k+(k−1)(d−1)=k(1+d−1)−1\text{感受野} = k + (k-1)(d-1) = k(1+d-1) - 1 感受野=k+(k−1)(d−1)=k(1+d−1)−1

dilation有效核3×3 的感受野
13×33
23×3(间隔1)5
43×3(间隔2)9

在保持分辨率的同时扩大感受野——这是分割任务的关键技术(DeepLab 系列的核心)。


四、1×1 卷积的特殊地位

4.1 它做什么

1×11 \times 11×1 卷积在每个空间位置上做跨通道的线性变换:

out[i,j,c]=∑c′Wc,c′⋅input[i,j,c′]\text{out}[i,j,c] = \sum_{c'} W_{c,c'} \cdot \text{input}[i,j,c'] out[i,j,c]=c′∑​Wc,c′​⋅input[i,j,c′]

空间维度不变,只混合通道。

4.2 两个关键用途

用途 1:升维 / 降维

python
nn.Conv2d(64, 128, kernel_size=1)   # 通道 64 → 128,不改变 H、W
nn.Conv2d(128, 64, kernel_size=1)   # 通道 128 → 64(降维)
1
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 为例):
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54

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

输出:

=== 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 的感受野: 11
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16

关键观察:stride=2 之后,跳跃间隔变成 2,后面每层的感受野增长更快(每次+2 而非 +1)。这是 stride 通过感受野公式的 jjj 项放大了后续所有层的效果。

实验 2:卷积 vs 全连接的参数量

python
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} 倍")
1
2
3
4
5
6
7
全连接:      7,554,585,344 参数
卷积:              1,792 参数
差距:          4,215,952 倍
1
2
3

这就是卷积统治视觉领域的直接原因:不是效果更好,而是能训得动。

实验 3:1×1 卷积做瓶颈

python
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}")
1
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])
1
2
3
4
5

实验 4:空洞卷积扩大感受野

python
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)的核心技巧:不下采样也能看得更远")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15

实验 5:卷积 vs 全连接的表达能力

python
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}")
1
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→2632 \to 30 \to 28 \to 26 32→30→28→26

输出 26×26

感受野:1+3×(3−1)=71 + 3 \times (3-1) = 71+3×(3−1)=7。

验证:输出是 26×26,说明需要从原图看 7×7 的区域——26+6=3226 + 6 = 3226+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→CoutC_{in} \to C_{out}Cin​→Cout​ 的线性变换。所以它能:

  1. 升降维:1×1 (256→64) 就是把每个像素的 256 维特征压到 64 维
  2. 混合通道:让不同输入通道的信息相互交流
  3. 构成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 是不可学习的固定操作,有两个问题:

  1. 信息损失:窗口内只保留最大值,其余信息全丢。且反向传播时梯度只路由到最大那个位置,其余位置梯度为 0——训练效率低
  2. 无法适应数据:不管输入是什么,下采样方式都一样

stride=2 卷积的优势:

  1. 可学习:下采样方式由数据决定
  2. 信息保留:是线性变换而非丢弃
  3. 能配合 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(更高!)← 加深反而更差
1
2

这不能用过拟合解释——如果是过拟合,训练误差应该很低但测试误差高。实际上 56 层的训练误差更高。

1.2 名字的由来

退化(degradation)= 不是过拟合,而是模型「变差」了。

更深的网络理论上应该能模拟浅层网络(后面几层学成恒等映射 y=xy = xy=x 就行)。但实验表明优化器找不到这个解。

1.3 一个关键实验(He 的诊断)

如果问题是「找不到恒等映射」,那构造一个「恒等捷径」应该有帮助:

python
# 在浅层网络旁边手工加一条恒等通路
out = F(x) + x    # F 网络在初始化时输出为 0,网络就等价于恒等映射
1
2

结果:56 层带捷径的网络,效果和 20 层一样好(没有退化)。

结论:退化不是表达能力问题,是优化问题——优化器难以在深层网络里找到「接近恒等」的解。


二、残差连接:形式化

2.1 结构变化

普通:y=F(x)y = F(x)y=F(x)

残差:y=F(x)+x\boxed{y = F(x) + x}y=F(x)+x​

反向传播时:

∂y∂x=∂F∂x+I\frac{\partial y}{\partial x} = \frac{\partial F}{\partial x} + I ∂x∂y​=∂x∂F​+I

那个 III(单位矩阵)是关键——梯度有一条恒为1 的直达通路。

2.2 为什么「恒等映射容易学」

优化器最难学的函数之一是「什么都不做」(F(x)=0F(x) = 0F(x)=0)。

  • 普通连接:y=F(x)y = F(x)y=F(x)。要让 y=xy = xy=x,需要 FFF 学一个恒等函数。ReLU 网络里学恒等很困难(需要正斜率权重穿过所有层)
  • 残差连接:y=F(x)+xy = F(x) + xy=F(x)+x。要让 y=xy = xy=x,只需 FFF 输出 0。输出 0 是最简单的解——把最后一层的权重和 bias 初始化为 0 即可

这就是残差连接的全部魔法:把「学恒等映射」这个难题变成了「什么都不学」。

2.3 ResNet block 的两种形式

标准 block(18/34层用):

python
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)  # ★ 加完再做 ReLU
1
2
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层用):

python
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)
1
2
3
4
5
6
7
8

为什么叫瓶颈:因为中间通道数只有外部的 1/4,形状像瓶颈。作用是降低计算量(见第 7 篇)。

2.4 维度不匹配怎么办

如果 xxx 和 F(x)F(x)F(x) 形状不同(比如通道数翻倍 + stride=2),不能直接相加。用一个 1×1 卷积做投影:

y=F(x)+W1×1∗xy = F(x) + W_{1\times1} * x y=F(x)+W1×1​∗x

注意:这个 projection 也是残差的一部分,梯度照样能通过。


三、残差连接的梯度分析(数学证明)

这是本文最有价值的部分。

3.1 梯度连乘展开

考虑 LLL 层残差网络,xl+1=xl+F(xl)x_{l+1} = x_l + F(x_l)xl+1​=xl​+F(xl​)。

反向传播:

∂L∂x1=∂L∂xL∏l=L−11∂xl+1∂xl=∂L∂xL∏l=L−11(I+∂Fl∂xl)\frac{\partial L}{\partial x_1} = \frac{\partial L}{\partial x_L}\prod_{l=L-1}^{1}\frac{\partial x_{l+1}}{\partial x_l} = \frac{\partial L}{\partial x_L}\prod_{l=L-1}^{1}\left(I + \frac{\partial F_l}{\partial x_l}\right) ∂x1​∂L​=∂xL​∂L​l=L−1∏1​∂xl​∂xl+1​​=∂xL​∂L​l=L−1∏1​(I+∂xl​∂Fl​​)

把这个乘积展开(关键):

∏l=1L(I+Jl)=I+∑lJl+∑l<kJlJk+⋯+∏lJl\prod_{l=1}^{L}\left(I + J_l\right) = I + \sum_l J_l + \sum_{l<k}J_lJ_k + \cdots + \prod_l J_l l=1∏L​(I+Jl​)=I+l∑​Jl​+l<k∑​Jl​Jk​+⋯+l∏​Jl​

所有项里,第一项就是 III。即使所有 JlJ_lJl​ 都趋近于 0(FFF 学成恒等映射),乘积也至少是 III——梯度完全不衰减。

3.2 一个更直观的理解

对比两种网络的梯度传递:

普通网络:

∂L∂x1=∂L∂xL∏lJl\frac{\partial L}{\partial x_1} = \frac{\partial L}{\partial x_L}\prod_l J_l ∂x1​∂L​=∂xL​∂L​l∏​Jl​

每个因子都要贡献自己的值。任何一个因子小,整体就衰减。

残差网络:

∂L∂x1=∂L∂xL(I+∑lJl+∑JlJk+⋯ )\frac{\partial L}{\partial x_1} = \frac{\partial L}{\partial x_L}\left(I + \sum_l J_l + \sum J_lJ_k + \cdots\right) ∂x1​∂L​=∂xL​∂L​(I+l∑​Jl​+∑Jl​Jk​+⋯)

「什么都不学」(所有 Jl=0J_l = 0Jl​=0)时,梯度是 ∂L∂xL\frac{\partial L}{\partial x_L}∂xL​∂L​ 原封不动地传下去。

3.3 残差块还能做什么

不只是「什么都不做」——如果需要修改表示:

  • F(x)=xF(x) = xF(x)=x → 恒等,保持信息
  • F(x)=0F(x) = 0F(x)=0 → 什么都不做(网络自己选的)
  • F(x)=F(x) = F(x)= 任意变换 → 加上新信息

关键洞察:残差块的表达能力是「xxx 加上任意函数」,不是「任意函数」。所以恒等映射永远在假设空间里,无论网络多深都不会丢失。


四、DenseNet:另一个思路

4.1 结构差异

ResNetDenseNet
连接方式xl+1=xl+Fl(xl)x_{l+1} = x_l + F_l(x_l)xl+1​=xl​+Fl​(xl​)(相加)xl+1=Cat([xl,Fl(xl)])x_{l+1} = \text{Cat}([x_l, F_l(x_l)])xl+1​=Cat([xl​,Fl​(xl​)])(拼接)
各层特征只传到下一层每层都传到所有后续层
参数利用后期层拿不到前期层的「原始特征」每层都能直接用所有前期特征

DenseNet 里的「稠密连接」:

x1 ──────────────────────────────┐
 ↓                                │
x2 = Cat([x1, F1(x1)]) ───────────┤
 ↓                                │
x3 = Cat([x2, F2(x2)]) ───────────┤   ← x1 的原始特征一直在
 ↓                ││
x4 = Cat([x3, F3(x3)]) ───────────┤
1
2
3
4
5
6
7

4.2 两者的对比

ResNet 的类比:残差连接是「在高速公路上加一个出口」——你可以选择走新路或者留在高速上。

DenseNet 的类比:DenseNet 是「每个景点都直达所有景点」——不需要绕路回去看之前看到的东西。

4.3 各自的代价

ResNetDenseNet
参数量较少通道数递增,后期层很大
显存占用较低高(要保存所有中间特征)
训练速度快慢
精度高略高(但现在 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 的思想改造 CNN
1
2
3
4
5
6
7
8
9
10
11
12
13

两条线索交织:

  1. 深度(VGG → ResNet):靠残差突破
  2. 注意力/全局视野(SE → ViT → Swin → ConvNeXt):从局部卷积走向全局

ConvNeXt 值得单独说:它证明了「Transformer 的设计原则(LayerNorm、大核、GELU、AdamW)」比「Transformer 的架构」更本质。理解 CNN 和 Transformer 孰优孰劣,比记住某个具体架构更有价值。


六、动手实验

实验 1:复现退化现象(普通网络加深反而变差)

这是本文最重要的实验。用一个需要深层非线性的目标函数(3 层 tanh 嵌套 + 线性),扫描网络深度。

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

实测输出:

深度    普通网络测试误差
    4         0.14639
   20         0.21508
   50         0.21671
  100         0.21673
1
2
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:残差块的梯度分析

验证「恒等通路」的梯度贡献。

python
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}")
1
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+05
1
2
3
4
5

这个对比极其震撼:

深度无残差有残差倍数
21.13e+014.82e+014×
109.19e-031.85e+022万倍
307.09e-112.24e+033×10¹³ 倍
600.000e+001.01e+05∞

60 层无残差网络,梯度精确变成 0 —— float32 已经表示不出这个数了,输入端的神经元完全收不到梯度,等价于一个随机初始化的浅网络。

60 层残差网络,梯度是 1.01e+05,健壮可用。

这就是「恒等通路 III」的数学保证:即使所有 FlF_lFl​ 学成恒等映射(Jl=0J_l = 0Jl​=0),梯度也至少原封不动地传下去。

这比实验 1 更能说明残差连接的本质——它不只是「让深网络能训」,而是让深网络的梯度完全健康。

实验 3:残差块的初始化技巧(让块初始就是恒等映射)

零初始化最后一层,则 F(x)=0F(x) = 0F(x)=0,整个残差块初始时就是恒等映射 y=xy = xy=x。

python
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())} 个")
1
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!
输出里有负数吗: 248
1
2
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 篇。

另一个要点:把最后一层归零后,训练初期整个网络等价于一堆恒等映射,非常稳定;随着训练进行,FFF 逐渐学到东西。这就是「让深网络从简单解开始学」的具体实现(这是 ControlNet 的核心技巧)。

七、自测题

Q1:如果把 ResNet block 里的 F(x) + x 改成 F(relu(x)) + x 会怎样?

答案

能用,但会损失一部分「恒等通路的干净性」。

具体来说:如果 xxx 里有负分量,relu(x) 会把它们变成 0,于是

out=F(relu(x))+x\text{out} = F(\text{relu}(x)) + x out=F(relu(x))+x

注意这里加的还是原始 xxx,所以:

  • 恒等通路 x→+→outx \to + \to \text{out}x→+→out 依然存在且未被破坏
  • 梯度仍能通过 + 1+\,1+1 直达

所以梯度流的保证还在。

真正的破坏发生在另一个变体:out = F(relu(x) + x)(先 ReLU 再残差)——这样输出全是非负,下一层的负信息丢了。

「先加后 ReLU」才是 ResNet 的正确做法:

python
out = conv2(...)      # 无 ReLU
out = out + identity  # ★ 先相加
out = relu(out)       # ★ 再 ReLU
1
2
3

如果先 ReLU 再相加,输出会有负值,恒等映射就不再是「什么都不做」了。

Q2:为什么 ResNet 能在很深的网络里工作,但把网络加深后理论表达能力并没有变强?

答案

因为残差块的表达能力上限是「xxx 加上任意函数」,而深层网络能做的组合复杂度增长远慢于层数的增长。

更具体地说:

  1. 理论上加深网络确实表达能力更强(万能逼近定理),但**「更强」和「能不能训出来」是两件事**
  2. ResNet 的贡献主要是优化友好性,不是表达能力的提升
  3. 理论上层数翻倍能表达的函数种类远超需要,但优化器找不到那个解

启发:深度带来的是「更容易找到好解」,而不是「能表示更多函数」。

论文里的实验支持:ResNet-200 的层数约为 ResNet-50 的 4 倍,但两者精度几乎一样(甚至略低)——说明在这个任务上,表达能力早已不是瓶颈。

这不是深度无用,而是「深度需要正确的架构来兑现」。

Q3:DenseNet 和 ResNet 的核心区别是什么?各自适合什么场景?

答案

核心区别:特征传递方式。

ResNetDenseNet
方式相加 xl+1=xl+Fl(xl)x_{l+1} = x_l + F_l(x_l)xl+1​=xl​+Fl​(xl​)拼接 xl+1=Cat([x0,...,xl])x_{l+1} = \text{Cat}([x_0,...,x_l])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×37×7(扩大感受野)
激活函数ReLUGELU(平滑)
缩放层conv+BN+ReLU ×21×1 conv + GELU + 1×1 conv(MLP 风格)
优化器SGD + weight decayAdamW + 0.05 wd

核心洞察:决定性能的不是「Transformer 这个架构」,而是它的设计原则:

  1. Pre-LN 结构稳定深层训练(第 6 篇讲的残差通路)
  2. 大核提供全局视野(卷积的等价替代品)
  3. MLP 式的通道混合(比堆 3×3 更有效的通道交互)

为什么能行:ViT 的优势来自「全局注意力 + 大感受野 + 深层」,前两个可以用卷积近似,第三个靠架构改进。如果你不需要真正的「动态权重」(attention 的 QKTQK^TQKT),卷积 + 大核 + Pre-LN 是一个更高效的替代方案。

工程收益:不需要 CUDA 相关的自定义算子,标准 PyTorch 就能跑,速度和硬件兼容性都更好。


下一篇 → 初始化与数值稳定

上一级: 目录 · 上一篇

9 · 初始化与数值稳定

核心问题:为什么不能用 0 初始化?Xavier 和 He 的公式怎么来的?混合精度训练为什么需要 loss scaling?


一、为什么初始化如此重要

初始化做两件事:

  1. 打破对称性——否则所有神经元学到的完全一样(见自测题)
  2. 控制激活值的方差在层间传播时保持恒定

第二点是关键。考虑一个 nin→noutn_{in} \to n_{out}nin​→nout​ 的线性层:

zj=∑i=1ninwjixi+bjz_j = \sum_{i=1}^{n_{in}} w_{ji} x_i + b_j zj​=i=1∑nin​​wji​xi​+bj​

如果 xxx 的方差是 σx2\sigma_x^2σx2​,权重独立同分布、方差 σw2\sigma_w^2σw2​,那么:

Var[zj]=nin⋅σw2⋅σx2\text{Var}[z_j] = n_{in} \cdot \sigma_w^2 \cdot \sigma_x^2 Var[zj​]=nin​⋅σw2​⋅σx2​

要让输出方差 = 输入方差(方差保持),需要:

nin⋅σw2=1⇒σw2=1nin\boxed{n_{in} \cdot \sigma_w^2 = 1 \quad \Rightarrow \quad \sigma_w^2 = \frac{1}{n_{in}}} nin​⋅σw2​=1⇒σw2​=nin​1​​

关键洞察:方差保持的条件只依赖于 ninn_{in}nin​,与激活函数无关。 但ReLU 会砍掉一半的激活,所以要补偿。


二、Xavier(Glorot)初始化

2.1 公式

σw2=2nin+nout\boxed{\sigma_w^2 = \frac{2}{n_{in} + n_{out}}} σw2​=nin​+nout​2​​

推导:同时让前向传播的方差保持和反向传播的梯度方差保持。

  • 前向(方差保持):Var[z]=ninσw2σx2\text{Var}[z] = n_{in}\sigma_w^2\sigma_x^2Var[z]=nin​σw2​σx2​,要等于 σx2\sigma_x^2σx2​ → σw2=1/nin\sigma_w^2 = 1/n_{in}σw2​=1/nin​
  • 反向(梯度保持):同理 σw2=1/nout\sigma_w^2 = 1/n_{out}σw2​=1/nout​
  • 兼顾两者:σw2=2nin+nout\sigma_w^2 = \frac{2}{n_{in}+n_{out}}σw2​=nin​+nout​2​

2.2 适用场景

Xavier 适用于 tanh/sigmoid 等关于原点对称的激活函数。

因为对称激活的正负部分都有贡献,不需要补偿。

2.3 局限

用 ReLU 时,Xavier 会让激活方差逐层减半。

因为 ReLU 砍掉负半轴,只剩一半的激活有贡献:

Var[a]=12Var[z](ReLU 后)\text{Var}[a] = \frac{1}{2}\text{Var}[z] \quad(\text{ReLU 后}) Var[a]=21​Var[z](ReLU 后)

而 Xavier 假设了「全部激活都有贡献」,所以实际方差会每层乘以 0.5。20 层后就是 0.520≈10−60.5^{20} \approx 10^{-6}0.520≈10−6——信号基本消失。


三、He(Kaiming)初始化

3.1 公式

σw2=2nin\boxed{\sigma_w^2 = \frac{2}{n_{in}}} σw2​=nin​2​​

推导:Xavier 基础上补偿 ReLU 砍掉一半:

2nin+nout≈22nin=1nin(当 nin≈nout)\frac{2}{n_{in}+n_{out}} \approx \frac{2}{2n_{in}} = \frac{1}{n_{in}} \quad (\text{当 } n_{in}\approx n_{out}) nin​+nout​2​≈2nin​2​=nin​1​(当 nin​≈nout​)

而我们要的是 2nin\frac{2}{n_{in}}nin​2​——正好是 2 倍。这 2 倍就是 ReLU 的补偿。

通用形式(PyTorch 的 nonlinearity 参数):

σ=2nin(1−p2)\sigma = \sqrt{\frac{2}{n_{in}\left(1 - p^2\right)}} σ=nin​(1−p2)2​​

其中 ppp 是 dropout 概率(He 论文考虑了 dropout 的影响)。

3.2 实测对比(这是本文最重要的实验)

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

实测输出:

     初始化     首层激活方差   末层激活方差        衰减倍数
     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+14x
1
2
3
4

这张表说明了一切:

初始化20 层后方差衰减判断
默认(uniform ±1/√fan_in)4.65e+14 倍数值下溢,完全失效
Xavier5.87e+05 倍仍然衰减太多
He2.62 倍✅ 基本恒定

He 初始化下,激活方差穿过 20 层只衰减 2.6 倍——这就是「方差保持」的定量证明。

而默认初始化衰减 4.65e14 倍,float32 早就溢出/下溢了。这就是为什么 PyTorch 早期版本在深网上训不起来。


四、PyTorch 的默认初始化

4.1 各层的默认值

层默认初始化说明
nn.LinearU(−1nin,1nin)U(-\frac{1}{\sqrt{n_{in}}}, \frac{1}{\sqrt{n_{in}}})U(−nin​​1​,nin​​1​)Kaiming uniform
nn.Conv2dKaiming uniform(按 fan_in)
nn.BatchNorm2dweight=1, bias=0恒等变换
nn.LayerNormweight=1, bias=0恒等变换
残差块最后一层可能 zeros见第 8 篇
EmbeddingN(0,1)N(0,1)N(0,1)LLaMA 用此初始化

注意 Linear 的默认其实是 Kaiming uniform((a=\sqrt{5}) 的变体),不是 Xavier。 但 PyTorch 的实现里 gain 算的是 1/fanin1/\sqrt{fan_{in}}1/fanin​​ 而不是 2/fanin\sqrt{2/fan_{in}}2/fanin​​——所以它严格来说既不是 He 也不是 Xavier,是一个偏保守的选择。

所以 PyTorch 里想用 He 初始化必须手动指定:

python
for m in model.modules():
    if isinstance(m, nn.Linear):
        nn.init.kaiming_normal_(m.weight, nonlinearity='relu')
        nn.init.zeros_(m.bias)
1
2
3
4

4.2 现代实践:为什么 LLM 全用 N(0,0.02)N(0, 0.02)N(0,0.02)

LLaMA / GPT 系列对所有线性层用同一个简单初始化:

python
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)
1
2
3
4
5
6

为什么可以这么简单粗暴? 三个原因:

  1. RMSNorm 在每个子层入口就把激活归一化了,所以进入下一个线性层时方差已经被控制,不需要针对每个层算 fan
  2. 残差连接 + Pre-LN 让梯度稳定,对初始化的敏感度降低
  3. 所有层形状相似(hidden_dim 统一),一个 std 够了

这是一个「架构设计降低了调参需求」的典型例子——因为 RMSNorm + 残差已经把问题解决了,初始化只需要「别太大就行」。


五、数值稳定性:FP16 与混合精度

5.1 问题的来源

fp16 的表示范围:10−5∼10510^{-5} \sim 10^510−5∼105。

反向传播时梯度会比激活值小几个数量级,容易下溢成 0:

fp32: 梯度 1e-8   → 正常
fp16: 梯度 1e-8   → 下溢成 0(fp16 最小正规数约 6e-5)
1
2

结果:深层模型的梯度全部消失,训练完全失效。

5.2 解决方案:loss scaling

python
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()                      # ★ 动态调整 S
1
2
3
4
5
6
7
8
9
10
11

机制:

实际梯度=S⋅真实梯度\text{实际梯度} = S \cdot \text{真实梯度} 实际梯度=S⋅真实梯度

把梯度放大 65536 倍,让它在 fp16 范围里可表示;scaler.step() 内部会除以 S 再更新参数。

动态调整 S:如果检测到梯度有 inf/nan,就减小 S(这次更新跳过);连续 N 次正常就增大 S。这样就不用手动猜缩放系数。

5.3 bf16:更简单的替代方案

bf16(bfloat16)用和 fp32 相同的指数位(8 位),只减少尾数位(7 位 vs 10 位)。

类型位分配范围精度
fp321+8+2310±3810^{\pm38}10±38高
fp161+5+1010±510^{\pm5}10±5中
bf161+8+710±3810^{\pm38}10±38低
fp81+4+310±210^{\pm2}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)=log⁡∑iezi=zmax⁡+log⁡∑iezi−zmax⁡\text{logsumexp}(z) = \log\sum_i e^{z_i} = z_{\max} + \log\sum_i e^{z_i - z_{\max}} logsumexp(z)=logi∑​ezi​=zmax​+logi∑​ezi​−zmax​

减最大值保证指数部分不会溢出。所以直接传 logits 是安全的(第 1 篇已详述)。

6.2 梯度的 NaN 排查

训练出现 nan 时的排查顺序:

python
# 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'])
1
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/inf8%清洗数据
梯度爆炸2%梯度裁剪

七、动手实验

实验 1:初始化的影响(本文核心实验)

见第二节的代码。必须自己跑一遍看那张表,三个数字的对比非常有说服力。

实验 2:对称性——为什么不能全零初始化

python
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("→ 三个输出不同,反向传播的梯度也不同,神经元才能分化")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26

实测输出:

输入: [[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']
→ 三个输出不同,反向传播的梯度也不同,神经元才能分化
1
2
3
4
5
6
7
8
9
10

坑:我第一版只把 weight 归零,忘了nn.Linear 默认有 bias(初始化为均匀随机值)。结果输出不全是 0,看起来「对称性问题不存在」——其实是 bias 打乱了假象。要做对称性实验,所有参数都要归零。

这是初学者最容易忽略但又最重要的一点:初始化不只是「让训练稳定」,它决定了各个神经元能不能分化出不同的功能。

实验 3:bf16 vs fp16 的下溢差异

python
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 同量级")
1
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 同量级
1
2
3
4
5
6
7
8

这就是 bf16 不需要 loss scaling 的原因:它的动态范围和 fp32 一样,不会下溢。

实验 4:loss scaling 的作用

python
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}")
1
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]≈2nin+nout\text{Var}[z] \approx \frac{2}{n_{in}+n_{out}}Var[z]≈nin​+nout​2​ 合理
  • ReLU:砍掉负半轴,只有一半的激活被保留

具体来说,Xavier 让 σw2=2nin+nout\sigma_w^2 = \frac{2}{n_{in}+n_{out}}σw2​=nin​+nout​2​。当 nin=nout=nn_{in}=n_{out}=nnin​=nout​=n 时,σw2=1n\sigma_w^2 = \frac{1}{n}σw2​=n1​。

前向传播:Var[z]=n⋅1n⋅Var[x]=Var[x]\text{Var}[z] = n \cdot \frac{1}{n} \cdot \text{Var}[x] = \text{Var}[x]Var[z]=n⋅n1​⋅Var[x]=Var[x] ✓ 保持

但 ReLU 之后:Var[a]=12Var[z]=12Var[x]\text{Var}[a] = \frac{1}{2}\text{Var}[z] = \frac{1}{2}\text{Var}[x]Var[a]=21​Var[z]=21​Var[x] ✗ 每层减半

20 层后:1220≈10−6\frac{1}{2^{20}} \approx 10^{-6}2201​≈10−6。实测衰减 5.87e5倍(本篇实验数据)。

He 的修正:把方差翻倍,σw2=2nin\sigma_w^2 = \frac{2}{n_{in}}σw2​=nin​2​,补偿 ReLU 砍掉的那一半。实测 20 层只衰减 2.62 倍。

选择规则:

  • ReLU / LeakyReLU / GELU → He (Kaiming)
  • tanh / sigmoid → Xavier (Glorot)

Q2:残差网络里,最后一层的 BN 为什么常初始化为 gamma=0?

答案

为了让残差分支在训练初期输出 0,即 F(x)=0F(x) = 0F(x)=0,此时整个残差块等价于恒等映射。

回忆 ResNet block 的结构:

out = conv2(bn2(conv1(bn1(x))))
identity = x(或者下采样后的 x)
out = relu(out + identity)
1
2
3

如果 BN 的 γ=0\gamma = 0γ=0,则 bn2 输出全 0 → out = 0 → relu(0 + identity)。

好处:

  1. 训练初期整个网络等价于恒等堆叠,非常稳定
  2. 残差分支「从零开始学」,而不是一开始就引入随机扰动
  3. 这是「让深网络从简单解开始学」的具体实现

同样的技巧用于 DiT / ControlNet:把最后一个线性层的权重和偏置初始化为 0,让网络初始时是「什么都不做」,然后逐渐学到有用的变换。

实测参考(本篇实验 3):零初始化后差异是 2.582.582.58,不是 0——因为末尾的 ReLU 会截断负数。所以严格来说不是精确恒等,但仍然大大稳定了训练。

Q3:混合精度训练时,为什么 scaler.unscale_(optimizer) 必须在 clip_grad_norm_ 之前?

答案

因为 clip_grad_norm_ 计算的是梯度范数,梯度此时被放大了 S 倍。

python
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)   # 现在范数是真实的
1
2
3
4
5
6
7
8
9
10
11

PyTorch 的 API 设计:scaler.step(optimizer) 内部会自动 unscale_,所以如果不用梯度裁剪,可以不手动调用。

但如果你手动裁剪,必须自己先 unscale_,否则裁剪的是错误的范数。

推荐的安全写法:

python
scaler.scale(loss).backward()
scaler.unscale_(optimizer)                  # 总是先 unscale
torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)
scaler.step(optimizer)
scaler.update()
1
2
3
4
5

Q4:为什么 bf16 不需要 loss scaling,但仍常配合 FP32 累积?

答案

不需要 loss scaling 的原因:bf16 的指数位和 fp32 相同(8 位),动态范围都是 10±3810^{\pm38}10±38。而梯度小到 10−810^{-8}10−8 时 fp16 会下溢,bf16 不会(本篇实验 3 实测:fp16 → 0,bf16 → 9.9e-8)。

仍然需要 FP32 累积的原因:bf16 只有 7 位尾数(fp32 是 23 位),精度很低。

具体问题:

  1. 累加误差:优化器更新参数时是 w←w−lr⋅gw \leftarrow w - lr \cdot gw←w−lr⋅g。当 lrlrlr 很小时(比如 10−810^{-8}10−8),lr⋅glr \cdot glr⋅g 相对于 www 极小,bf16 的 7 位尾数根本无法表示这个微小变化,更新会被完全舍入丢弃。
  2. 梯度累加时的抵消:小梯度的累加容易产生灾难性抵消。

标准做法:

python
# 前向 + 反向: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 只影响前向的计算精度
1
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 年的主流设计。

#章节核心问题
10Attention 的数学推导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 篇里的每个组件在这份代码里都有对应实现。

bash
pip install torch
python llama_from_scratch.py    # 前向 + loss + 反向,全程通过
1
2

这一部分要建立的判断力

  • 看到一个 LLM 架构图,能逐个组件说出它为什么在那里
  • 能算清一次推理的显存占用和 FLOPs,判断瓶颈在哪
  • 能区分训练期优化和推理期优化 —— 两者的约束完全不同

10 · Attention 的数学推导

核心问题:QKV 为什么要除 d\sqrt{d}d​?softmax 为什么必须有?如果去掉 scaling 会怎样?


一、从「需求」出发

1.1 我们想要什么能力

处理序列时,模型需要能:让每个位置「主动去查」其他位置的相关信息。

比如处理「那只猫很可爱,因为它饿了」——"它"要能关联到前面的"猫"。这种关联是动态的、依赖内容的,不能靠固定位置编码。

1.2 用检索类比理解 QKV

把 Attention 想成一个数据库检索系统:

角色类比作用
Q(Query)我想要什么当前查询的「需求描述」
K(Key)我有什么每个条目的「索引标签」
V(Value)实际内容每个条目真正的数据

检索流程:

  1. 拿我的 Q 去和所有条目的 K 比对 → 得到匹配分数
  2. 把分数转成权重(softmax)→ 归一化到 [0,1]
  3. 按权重加权求和所有条目的 V → 得到我要的结果

关键洞察:相似度是「Q 和 K 的关系」,但真正被加权取回的是 V。Q/K 负责「找谁」,V 负责「拿什么」。 这就是为什么 V 通常不参与打分。


二、Scaled Dot-Product Attention 的公式

Attention(Q,K,V)=softmax(QK⊤dk+M)V\text{Attention}(Q,K,V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} + M\right) V Attention(Q,K,V)=softmax(dk​​QK⊤​+M)V

其中:

  • Q∈Rn×dkQ \in \mathbb{R}^{n \times d_k}Q∈Rn×dk​(nnn 个查询,每个 dkd_kdk​ 维)
  • K∈Rm×dkK \in \mathbb{R}^{m \times d_k}K∈Rm×dk​(mmm 个键)
  • V∈Rm×dvV \in \mathbb{R}^{m \times d_v}V∈Rm×dv​
  • MMM 是掩码矩阵(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]
1
2
3
4
5
6
7
8
9

注意中间那个 [n,m][n, m][n,m] 矩阵是核心:它就是「注意力图」,可视化出来能看到模型在学什么。


三、三个核心问题

3.1 为什么除 dk\sqrt{d_k}dk​​(这是本篇最重要的部分)

数学推导:

假设 qqq 和 kkk 的每个分量都是独立的、均值 0 方差 1 的随机变量。那么点积:

q⋅k=∑i=1dkqikiq \cdot k = \sum_{i=1}^{d_k} q_i k_i q⋅k=i=1∑dk​​qi​ki​

方差(独立项相加,方差相加):

Var[q⋅k]=∑i=1dkVar[qiki]=∑i=1dkVar[qi]Var[ki]=dk⋅1⋅1=dk\text{Var}[q \cdot k] = \sum_{i=1}^{d_k}\text{Var}[q_i k_i] = \sum_{i=1}^{d_k} \text{Var}[q_i]\text{Var}[k_i] = d_k \cdot 1 \cdot 1 = d_k Var[q⋅k]=i=1∑dk​​Var[qi​ki​]=i=1∑dk​​Var[qi​]Var[ki​]=dk​⋅1⋅1=dk​

所以 std[q⋅k]=dk\text{std}[q\cdot k] = \sqrt{d_k}std[q⋅k]=dk​​。

实测验证:

python
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}")
1
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.040
1
2
3
4

完美吻合:点积的 std 就是 d\sqrt{d}d​,除以 d\sqrt{d}d​ 后稳定在 1 附近。

为什么这会导致训练失败:

softmax 的输入被放大 d\sqrt{d}d​ 倍后进入饱和区。饱和的 softmax 是什么样子? 看实验:

python
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}")
1
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.545
1
2
3
4
5
6
7
8
9
10
11
12
13

仔细看 logits_std = 16 那一行:熵只有 0.105 ~ 0.326,而均匀分布的熵是 1.386 ~ 5.545。注意力分布几乎变成了 one-hot——每个查询只关注一个键。

这会导致三个后果:

  1. 梯度消失:softmax 饱和区域的导数 softmax(z)(1−softmax(z))\text{softmax}(z)(1-\text{softmax}(z))softmax(z)(1−softmax(z)) 趋近 0
  2. 无法学习:所有注意力集中在一个位置上,模型丧失了「对比多个位置」的能力
  3. 信息瓶颈:每个 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\sqrt{d_k}dk​​ 不是可选的微调,而是保证 softmax 工作在有效区间的必要操作。

3.2 为什么必须用 softmax(不能直接用点积)

三个理由:

理由1:需要归一化。softmax 让权重和为 1,等于「分配注意力预算」——每个 token 分配的注意力总量固定,不会因为某个分数高就无限放大。

理由2:可微且梯度有意义。softmax 是平滑函数,梯度 ∂ai∂sj=ai(δij−aj)\frac{\partial a_i}{\partial s_j} = a_i(\delta_{ij} - a_j)∂sj​∂ai​​=ai​(δij​−aj​) 给出「提升 key j 的分数会如何改变分配」的清晰信号。

理由3:引入竞争/相对性。softmax 是「相对」操作——某个键分数升高会抢占其他键的权重。这符合注意力的直觉:注意力是稀有的资源,需要竞争。

如果直接用点积(可以理解为加权平均但不归一化),权重可能全都很小(输出尺度失控)或都很大(输出爆炸)。

3.3 mask 的作用

softmax(QK⊤dk+M)\text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} + M\right) softmax(dk​​QK⊤​+M)

其中 MMM 在要屏蔽的位置是 −∞-\infty−∞(softmax 后变成 0)。

Causal mask(自回归模型):第 iii 个 token 只能看到 ≤i\le i≤i 的 token。

python
mask = torch.triu(torch.ones(seq, seq) * float('-inf'), diagonal=1)
1

用 -inf 而不是大负数的原因:softmax 会先减最大值,-inf 在减法后会变成 NaN(-inf−finite=-inf\text{-inf} - \text{finite} = \text{-inf}-inf−finite=-inf,exp⁡(−∞)=0\exp(-\infty) = 0exp(−∞)=0 实际是安全的,但减法顺序可能出问题)。实践中 PyTorch 用 float('-inf') 是安全的,因为 softmax 内部对 -inf 有特殊处理。


四、完整实现

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

PyTorch 内置版本(生产环境务必用这个):

python
# ★ 官方推荐:内存高效的融合实现
out = F.scaled_dot_product_attention(Q, K, V, is_causal=True)
1
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]
1
2
3
4
5
6
7
8
python
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]])
1
2
3
4
5
6
7

5.2 Padding mask 的组合

实际场景要同时考虑 causal mask 和 padding mask:

python
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, :]
1
2
3
4
5
6

六、动手实验

实验 1:验证 √d 缩放(本文核心)

python
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,有明确的相对差异可供学习")
1
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,有明确的相对差异可供学习
1
2
3
4
5
6
7
8
9

注意 max_p(缩放) 那一列只有 0.02~0.035,看起来「过于均匀」了——但这恰恰是正确的:真实训练中的 q,kq, kq,k 不是独立随机向量,它们经过投影和训练后有特定的相关结构,会让部分位置的分数更高。随机初始化下的均匀分布正是我们想要的起点(有区分度,才能学)。

真正要避免的是 max_p(不缩放) 那一列的 0.9999——所有位置的注意力完全相同,没有任何区分能力。

注意:真实 Transformer 里 q,kq, kq,k 经过 WQ,WKW_Q, W_KWQ​,WK​ 投影后分量 std 约 1/din1/\sqrt{d_{in}}1/din​​ 量级,但经过层归一化和训练后,实际进入 attention 的 q,kq, kq,k 分量 std 接近 1,所以上面的分析成立。

实验 2:手动实现并与 PyTorch 对比

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

实验 3:可视化注意力权重

python
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 列")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20

观察:第 4~7 行(副本)应该主要指向对应的 0~3 列,因为它们内容相同,q⋅kq\cdot kq⋅k 最大。这直观展示了 attention 在「按内容匹配」。


七、自测题

Q1:如果去掉 dk\sqrt{d_k}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⊤\text{diag}(p) - pp^\topdiag(p)−pp⊤,元素大小 ∼pi(1−pj)\sim p_i(1-p_j)∼pi​(1−pj​)。当 p→1p \to 1p→1 时 p(1−p)→0p(1-p) \to 0p(1−p)→0,梯度消失。

dk=1024d_k = 1024dk​=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⊤VQ^\top VQ⊤V 直接算?那得到的是 [d,d][d, d][d,d] 的矩阵(不是注意力权重矩阵),失去了「按需检索」的能力。

类比:图书馆系统里,「检索词→书的索引」和「书的实际内容」必须分开存储。如果混在一起,就无法实现「用检索词找到匹配的书,然后取它的内容」。

一个反例说明问题:假设一句话里有两个关键信息点,V 分别是它们的内容。如果只有一个矩阵,你无法表达「我要同时关注这两点,但注意力要按相关性分配」。

实际上 K 通常可以和 Q 共享同一份计算结果(self-attention 里),某些简化架构会这么做,但效果会下降。标准实现里 W_Q、W_K、V_W 是三套独立参数。

Q3:为什么 causal mask 要用 −∞-\infty−∞,用 −109-10^9−109 行不行?

答案

理论上都行,实践中都用 −∞-\infty−∞(或者 PyTorch 的 float('-inf'))。

用 −109-10^9−109 的情况:

softmax([5.0, -1e9]) = [1.0, 0.0]   ✓ 效果正确
1

因为 exp⁡(−109)≈0\exp(-10^9) \approx 0exp(−109)≈0,和 exp⁡(−∞)=0\exp(-\infty) = 0exp(−∞)=0 几乎没区别。

但有两个隐患:

  1. 数值精度:如果模型内部用 fp16,−109-10^9−109 直接超出 fp16 范围(最大 65504),会变成 -inf 或 nan。用 fp32 的话 −109-10^9−109 也在边缘。

  2. 和 causal mask 组合时:如果同时有 padding mask,两层 mask 相加可能得到 −2×109-2\times10^9−2×109,进一步溢出。

PyTorch 的 F.scaled_dot_product_attention 用 is_causal=True 参数,内部自动处理这个 mask,比手动传更高效(能融合进 kernel)。

一个实用建议:训练时优先用 is_causal=True 而不是手动构造 mask,能享受 Flash Attention 的优化。

Q4:注意力矩阵 [n,m][n, m][n,m] 的复杂度是多少?为什么这是长序列的瓶颈?

答案

时间复杂度:O(n⋅m⋅dk)O(n \cdot m \cdot d_k)O(n⋅m⋅dk​),当 n=m=Ln = m = Ln=m=L 时是 O(L2dk)O(L^2 d_k)O(L2dk​)。

空间复杂度:注意力矩阵本身要存 O(L2)O(L^2)O(L2)。这是平方级。

具体数字(LLaMA-7B,4096 token):

注意力矩阵:2 (batch) × 32 (heads) × 4096 × 4096 × 2 bytes (fp16)
          = 2 × 32 × 4096 × 4096 × 2
          = 2.1 GB   ← 单层!
1
2
3

32 层累计就是 68 GB 仅用于存注意力矩阵。

对比 MLP 的复杂度:O(L⋅d2)O(L \cdot d^2)O(L⋅d2),线性于 LLL。

所以瓶颈很明确:

组件复杂度随序列长度
QKV 投影O(Ld2)O(Ld^2)O(Ld2)线性
注意力矩阵O(L2d)O(L^2 d)O(L2d)平方 ⚠️
MLPO(Ld2)O(Ld^2)O(Ld2)线性

这就是所有「高效注意力」工作的动机:

  1. FlashAttention:不显式存注意力矩阵(见第 14 篇)
  2. 稀疏注意力:只算部分位置(Longformer、BigBird)
  3. 线性注意力:改变计算顺序,用 K⊤VK^\top VK⊤V 先算(Linear Transformer、Performer)
  4. 低秩近似:把 QK⊤QK^\topQK⊤ 近似成低秩矩阵(Linformer)

GPT-3 的选择:只支持 2048 token 上下文,因为再长平方成本太高。GPT-4 的 128K 上下文能实现,靠的是 FlashAttention + 多种优化(推测,多层特征)。


下一篇 → 多头注意力与位置编码

上一级: 目录 · 上一篇

11 · 多头注意力与位置编码

核心问题:多头到底在多做什么?位置信息怎么注入?RoPE 的旋转技巧是什么?


一、多头注意力:为什么要「多头」

1.1 单头的限制

单个注意力头只能算出一个 [n,n][n, n][n,n] 的注意力矩阵。这意味着每个 token 只能有一个「关注模式」。

问题:语言里的关联是多模态的。同一句话里可能同时需要:

  • 语法关联:「猫」↔「的」(结构)
  • 语义关联:「猫」↔「动物」(指代)
  • 语用关联:这个 token 在整体语境中的角色

一个头做不到兼顾。

1.2 多头的做法

把 dmodeld_{model}dmodel​ 拆成 hhh 份,每份独立做注意力,最后拼接。

MultiHead(Q,K,V)=Concat(head1,…,headh)WO\text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1,\dots,\text{head}_h)W^O MultiHead(Q,K,V)=Concat(head1​,…,headh​)WO

headi=Attention(QWiQ,KWiK,VWiV)\text{head}_i = \text{Attention}(QW_i^Q, KW_i^K, VW_i^V) 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]
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18

注意一个容易误解的点:多头不是把 dmodeld_{model}dmodel​ 分成 hhh 份各自独立算完就不管了,而是每个头有自己独立的 WQ,WK,WVW_Q, W_K, W_VWQ​,WK​,WV​,参数完全不共享。最后还要过一个 WOW^OWO 做融合。

参数量:

项目参数量
WQ,WK,WVW_Q, W_K, W_VWQ​,WK​,WV​3×dmodel×dmodel3 \times d_{model} \times d_{model}3×dmodel​×dmodel​
WOW^OWOdmodel×dmodeld_{model} \times d_{model}dmodel​×dmodel​
总计4dmodel24 d_{model}^24dmodel2​

和单头完全一样! 多头不增加参数量,只是改变了计算的方式(把一个大矩阵乘法拆成 hhh 个小的)。这是「结构先验」而非「容量增加」的典型例子。

1.3 参数量验证

python
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:,}  ← 和多头相同")
1
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.0x
1
2
3

这说明:多头和单头的参数量相同(都是 4 个 dmodel×dmodeld_{model}\times d_{model}dmodel​×dmodel​ 矩阵),因为单头也需要 WQ,WK,WV,WOW_Q, W_K, W_V, W^OWQ​,WK​,WV​,WO 四个矩阵。多头是「把一个大注意力拆成 h 个小注意力」,不是「多加了参数」。

1.4 为什么小头反而更好

一个反直觉但被广泛验证的结论:减小 dkd_kdk​ 往往提升效果。

模型dmodeld_{model}dmodel​头数 hhhdkd_kdk​
原始 Transformer512864
LLaMA-7B409632128
LLaMA-13B512040128
LLaMA-65B819264128
LLaMA2-7B409632128

关键观察:几乎所有 LLM 的 dkd_kdk​ 都固定在 64 或 128,不管模型多大。

原因(推测性的,但有实证支持):

  1. 注意力矩阵的噪声:dkd_kdk​ 越大,点积的方差越大,softmax 越容易进入需要精细 dk\sqrt{d_k}dk​​ 校准的区域
  2. 过拟合风险:dkd_kdk​ 越大,每个头的表达能力越强,越容易过拟合
  3. 「多而浅」优于「少而深」:更多头 = 更多并行的关系模式,每头更专注

这个观察的价值:现代 LLM 的设计里,头数随模型规模线性增长,但 dkd_kdk​ 恒定——这意味着 dmodel=h×dkd_{model} = h \times d_kdmodel​=h×dk​,即模型的宽度主要由「头数」决定。


二、位置编码:Transformer 的先天缺陷

2.1 问题

Attention 是置换等变的(permutation-equivariant):

Attention(PX,PK,PV)=P Attention(X,K,V)\text{Attention}(PX, PK, PV) = P\,\text{Attention}(X,K,V) 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,kq, kq,k 向量按位置「旋转」。

3.2 二维情况下的旋转(理解的关键)

假设 d=2d = 2d=2,把向量看作复平面上的一个点:

q=(q0,q1)↔q0+iq1q = (q_0, q_1) \leftrightarrow q_0 + i q_1 q=(q0​,q1​)↔q0​+iq1​

位置 mmm 对应一个旋转角度 mθm\thetamθ。旋转:

q′=Rmq=(cos⁡mθ−sin⁡mθsin⁡mθcos⁡mθ)(q0q1)q' = R_m q = \begin{pmatrix} \cos m\theta & -\sin m\theta \\ \sin m\theta & \cos m\theta \end{pmatrix}\begin{pmatrix} q_0 \\ q_1 \end{pmatrix} q′=Rm​q=(cosmθsinmθ​−sinmθcosmθ​)(q0​q1​​)

为什么这个设计是天才的?

核心性质:旋转内积定理——

⟨Rmq,Rnk⟩=⟨q,Rn−mk⟩=f(q,k,n−m)\boxed{\langle R_m q, R_n k\rangle = \langle q, R_{n-m}k\rangle = f(q, k, n - m)} ⟨Rm​q,Rn​k⟩=⟨q,Rn−m​k⟩=f(q,k,n−m)​

内积只依赖相对距离 n−mn - mn−m,与绝对位置无关!

这意味着 RoPE 天然编码了相对位置关系——而这正是注意力真正需要的信息(「这个词和前一个词的关系」比「这个词在第 500 位」更有用)。

3.3 推广到高维

ddd 维向量切成 d/2d/2d/2 对,每对用不同频率的旋转:

θi=base−2i/d,i=0,1,…,d/2−1\theta_i = \text{base}^{-2i/d}, \qquad i = 0,1,\dots,d/2-1 θi​=base−2i/d,i=0,1,…,d/2−1

几何间隔的频率:θ\thetaθ 按指数衰减,所以低维用高频(捕捉近距离),高维用低频(捕捉远距离)。

python
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) # 旋转第二半
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18

3.4 RoPE 的两个关键性质(实测验证)

python
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("→ 不同距离内积不同 → 模型能区分相对距离")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25

实测输出:

同一个 (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-06
1
2
3
4
5
6
7
8

完美的验证:每个距离下的内积完全一致(波动是浮点误差级别),不同距离内积不同。

这正是 RoPE 优于绝对位置编码的核心原因:绝对位置编码下,同一对 (q,k) 在不同绝对位置的相似度不同,模型要额外学习「位置偏移」;RoPE 直接把这个关系做进了旋转里。

3.5 RoPE 的外推能力

RoPE 的外推能力来自频率的连续性:训练时见过的最大位置是 LLL,理论上可以推广到任意位置(只要角度还在有效范围内)。

但实际上会退化。原因是注意力熵爆炸:

  • 训练时模型学会了「近处的 token 重要」
  • 位置超过训练长度后,旋转角度过大,注意力分布变得不稳定
  • 结果:注意力熵(不确定性)急剧上升,模型开始「乱看」

改进方案(NTK-aware scaling / YaRN):

方法思路
位置插值(PI)把位置 mmm 缩放到 mLtarget/Ltrain\frac{m}{L_{\text{target}}/L_{\text{train}}}Ltarget​/Ltrain​m​
NTK-aware修正 base 参数,低频维度保持、高频维度插值
YaRN分维度组合 PI 和 NTK,理论上最优

实践建议:扩展上下文长度时,第一选择总是 YaRN 或 NTK-aware,比直接 PI 效果好得多。


四、ALiBi:另一种思路

ALiBi(Attention with Linear Biases) 更简单:不做旋转,直接在 softmax 之前给 logits 加一个位置偏置。

Attention=softmax(QK⊤d+bias)\text{Attention} = \text{softmax}\left(\frac{QK^\top}{\sqrt{d}} + \text{bias}\right) Attention=softmax(d​QK⊤​+bias)

biasij={0j≤i−α(i−j)j>i\text{bias}_{ij} = \begin{cases} 0 & j \le i \\ -\alpha(i - j) & j > i \end{cases} biasij​={0−α(i−j)​j≤ij>i​

其中 α\alphaα 是每个头的斜率参数(如 8 个头的 α\alphaα = [1/21,1/22,…,1/28][1/2^1, 1/2^2, \dots, 1/2^8][1/21,1/22,…,1/28])。

特点:

  • 极简:不需要任何参数,不增加计算
  • 外推性好:距离是线性的,理论上可外推
  • 性能略低于 RoPE:现代 LLM 基本不用了

为什么 RoPE 更好:ALiBi 是「硬性惩罚远处 token」,而 RoPE 是「让模型自己学位置关系」。后者更灵活。


五、动手实验

实验 1:验证多头不增加参数量

见第一节代码。实测:单头和多头都是 4 个 dmodel2d_{model}^2dmodel2​ 矩阵,参数量相同。

实验 2:RoPE 的相对位置性质

见第三节代码。这是本篇核心实验,输出显示内积波动仅 1e-6。

实验 3:位置编码的必要性

python
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("→ 这就是为什么必须要有位置编码!")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27

实验 4:频率的指数结构

python
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("→ 几何级数设计让每个维度覆盖一个不同的距离尺度")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15

六、自测题

Q1:多头注意力的参数量和单头一样,那多头的「多」体现在哪里?

答案

参数量一样,但「计算方式的结构」不同。

单头:QK⊤QK^\topQK⊤ 是 [n,n][n, n][n,n] 的矩阵,每个元素是 dmodeld_{model}dmodel​ 维向量的内积。

多头:把 dmodeld_{model}dmodel​ 拆成 hhh 份,每个头用 dk=dmodel/hd_k = d_{model}/hdk​=dmodel​/h 维做内积,共 hhh 个 [n,n][n,n][n,n] 矩阵。

「多」体现在:

  1. hhh 个独立的注意力图——同一对 token 在不同头上可以有完全不同的关联强度。比如某个头关注语法邻接,另一个头关注长距离指代。
  2. 每个头学到了不同的关系模式(可解释性研究已证实:有些头专门做「前一个词」的关注,有些头做「句号」的关注)
  3. 子空间多样性——不同头在不同的特征子空间里工作,类似 CNN 的多通道

类比:单头像「用一副有色眼镜看世界」,多头像「戴 8 副不同的眼镜」,每副看到的东西不一样。

参数量不变是优势:可以在不加参数的情况下增加表达多样性。

Q2:RoPE 为什么比「加位置向量」更好?至少说两个理由。

答案

理由 1:内置了相对位置关系

  • 加位置向量:q=(WQxi+pi)q = (W_Q x_i + p_i)q=(WQ​xi​+pi​),相似度里有 pi⊤pjp_i^\top p_jpi⊤​pj​ 这种「绝对位置对绝对位置」的项,是混乱的
  • RoPE:⟨Rmq,Rnk⟩=f(q,k,n−m)\langle R_m q, R_n k\rangle = f(q,k,n-m)⟨Rm​q,Rn​k⟩=f(q,k,n−m),干净地只依赖相对位置

理由 2:不占据表示空间

加位置向量时,dmodeld_{model}dmodel​ 维里有一部分专门用来存位置信息,挤占了表示内容的空间。

RoPE 不增加任何维度——它只是把 q,kq, kq,k 旋转了一下,维度完全不变。

理由 3:更好的外推

几何级数的频率设计让 RoPE 在位置上更平滑。绝对位置嵌入(一个可学习的向量表)在超出训练长度时完全没有对应向量,直接失效。

理由 4:与 attention 的数学结构兼容

RoPE 的旋转是正交变换,保持内积关系(旋转矩阵满足 R⊤R=IR^\top R = IR⊤R=I)。这保证了它不会破坏 q,kq, kq,k 原本的语义结构。

一句话总结:RoPE 是「把位置编码进操作里」,而不是「把位置编码进数据里」。前者更优雅、更高效。

Q3:为什么 LLM 的 dkd_kdk​ 几乎都固定为 64 或 128,不管模型多大?

答案

这是实证观察(可以查 HuggingFace 的 config.json 验证),理论解释有几种推测:

推测 1:注意力熵与性能的关系

dkd_kdk​ 越大,q⋅kq\cdot kq⋅k 的方差越大。即使除以 dk\sqrt{d_k}dk​​,模型也需要精细校准才能让 softmax 工作在好的区间。适中的 dkd_kdk​ 让 attention 更容易学。

推测 2:过拟合

dkd_kdk​ 大 → 每个头的表达能力更强 → 更容易过拟合。LLM 靠数据规模控制过拟合,不需要靠小 dkd_kdk​。

推测 3:「多而浅」的归纳偏置

更多头 = 更多并行关系模式,每头更专注。这比「少而深」的单头更符合注意力机制的本质——它本来是做「关系匹配」的,不是做「特征提取」的。

推测 4:工程效率

小 dkd_kdk​ 的矩阵乘法更容易被 GPU 的 tensor core 优化(尤其是 fp16/bf16)。dk=128d_k=128dk​=128 是 tensor core 的友好尺寸。

实践建议:不要自己改这个。除非你有明确的理由和验证,否则跟随主流配置。

Q4:RoPE 能外推到比训练时更长的序列吗?会遇到什么问题?

答案

理论上能,实践中会退化。

问题:注意力熵爆炸

训练时模型学会了「近处重要、远处次要」的模式。位置超出训练长度后:

  1. 旋转角度过大——base=10000base=10000base=10000 下,高维的旋转周期很长,但低维在位置 10610^6106 时已经转了无数圈,数值上不再有意义
  2. 注意力分布变得均匀或混乱——模型「不知道该关注谁」
  3. 困惑度急剧上升——实测通常从 10飙到 100+

解决方案(按推荐度):

方法核心思路效果
YaRN分维度组合 PI + NTK★★★ 目前最好
NTK-aware scaling修正 base,高频少插值、低频多插值★★☆ 简单有效
位置插值(PI)位置 m→m⋅LtrainLtargetm \to m \cdot \frac{L_{\text{train}}}{L_{\text{target}}}m→m⋅Ltarget​Ltrain​​★★ 有损
继续预训练在长文本上继续训★★★ 慢但最可靠

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 层                   │
                                └──────────────────────┘
1
2
3
4
5
6
7
8
9
10

两个关键区别:

EncoderDecoder
注意力 mask无(双向)causal(只能看前面)
注意力输入只有源序列自注意力 + encoder 的输出
用途BERT(理解)GPT(生成)

现代 LLM 只保留了 Decoder 部分(叫 decoder-only),因为:

  1. 统一了理解和生成任务
  2. 因果 mask 让训练可以并行(每个位置同时预测下一个词)
  3. 架构更简单

1.2 原始 Block 的 Post-LN 结构

out=LayerNorm(x+Sublayer(x))\text{out} = \text{LayerNorm}(x + \text{Sublayer}(x)) out=LayerNorm(x+Sublayer(x))

注意 LayerNorm 在残差相加之后 —— 这叫 Post-LN。

问题:LN 夹在残差通路上,梯度要穿过 LN 才能回传,导致深层训练需要精细的 warmup,否则容易发散。


二、现代 LLaMA 风格的四大改动

Post-LN                →  Pre-LN
LayerNorm → RMSNorm    →  RMSNorm(更轻)
位置编码相加            →  RoPE(更自然)
FFN 的 ReLU           →  SwiGLU(更有效)
1
2
3
4

核心主题:从「能不能训」转向「怎么训得更好」。


三、Pre-LN:最重要的改动

3.1 结构对比

Post-LN (原始 Transformer):
    x → Sublayer(x) → +x → LayerNorm → 下一层

Pre-LN (现代 LLM):
    x → LayerNorm(x) → Sublayer → +x → 下一层
1
2
3
4
5

3.2 为什么 Pre-LN 更好

梯度通路的区别:

Post-LN 的残差通路:不干净——LN 在通路中间,梯度要「穿过 LN」:

∂xL+1∂xL=LN′(xL+F(xL))(I+F′(xL))\frac{\partial x_{L+1}}{\partial x_L} = \text{LN}'\left(x_L + F(x_L)\right)\left(I + F'(x_L)\right) ∂xL​∂xL+1​​=LN′(xL​+F(xL​))(I+F′(xL​))

那个 LN′\text{LN}'LN′ 是额外因子,深层累积会不稳定。

Pre-LN 的残差通路:干净的恒等映射:

∂xL+1∂xL=I+F′(LN(xL))\frac{\partial x_{L+1}}{\partial x_L} = I + F'(\text{LN}(x_L)) ∂xL​∂xL+1​​=I+F′(LN(xL​))

那个 III 保证了梯度 100% 直通——即使 F′F'F′ 是任何矩阵,梯度至少原封不动传下去。

3.3 实证

Post-LNPre-LN
深层训练需要 warmup稳定
warmup必需(通常 4000 步)可选
最终性能可能略好略差或持平
现代 LLM几乎不用全部采用

为什么最终性能反而 Post-LN 略好? 一个解释是 Post-LN 的 LN 在每个子层后重新缩放特征,有一定的正则化作用。但这个优势远小于「能稳定训练」的价值,所以全部转向 Pre-LN。

一个易错点:Pre-LN 要求最后有一个 final norm(在所有层之后):

python
x = block(x) for _ in range(N)
x = self.norm(x)       # ★ final LayerNorm/RMSNorm,必需!
1
2

因为 Pre-LN 的最后一个 block 输出没有经过归一化(LN 在子层内部)。


四、RMSNorm

第 6 篇已详述。这里只强调它在 LLaMA 里的作用:

LayerNormRMSNorm(LLaMA 选择)
参数2d2d2d(γ,β\gamma, \betaγ,β)ddd(只有 γ\gammaγ)
计算求均值 + 求方差 + 减均值 + 除只求均方根 + 除
能否融合 kernel部分可以完全融合

LLaMA-7B 的实际节省(d=4096d=4096d=4096,32 层):

  • 参数:每层省 409640964096,共省 131K131K131K(占总量 0.02%,微不足道)
  • 显存和速度:可融合进相邻算子,减少多次 HBM 读写

关键点:RMSNorm 的收益不在参数,而在计算融合。


五、SwiGLU:更好的 FFN

5.1 三种 FFN 对比

原始 Transformer:

FFN(x)=ReLU(xW1+b1)W2+b2\text{FFN}(x) = \text{ReLU}(xW_1 + b_1)W_2 + b_2 FFN(x)=ReLU(xW1​+b1​)W2​+b2​

GLU(Gated Linear Unit):

GLU(x)=ReLU(xW1)⊗(xW2)\text{GLU}(x) = \text{ReLU}(xW_1) \otimes (xW_2) GLU(x)=ReLU(xW1​)⊗(xW2​)

SwiGLU(LLaMA 的选择):

SwiGLU(x)=SiLU(xW1)⊗(xW3)\text{SwiGLU}(x) = \text{SiLU}(xW_1) \otimes (xW_3) SwiGLU(x)=SiLU(xW1​)⊗(xW3​)

其中 SiLU(也叫 Swish):SiLU(x)=x⋅σ(x)\text{SiLU}(x) = x \cdot \sigma(x)SiLU(x)=x⋅σ(x)

5.2 门控的直觉

乘性激活 = 一个「门」控制另一个「内容」:

SiLU(xW₁)⊗ (xW₃)
   ↓            ↓
  门(0~1)      内容(任意实数)
   ↓            ↓
   ↑──── 逐元素相乘 ────↑
       门=0 → 完全关闭
       门=1 → 完全通过
1
2
3
4
5
6
7

为什么这有用:普通的 ReLU 只能「开启或关闭」整个神经元(且不可微地调整)。门控允许每个维度独立控制信息的通过量,而且是可学习的。

5.3 参数量陷阱:为什么中间层要缩小

SwiGLU 有三个矩阵(W1,W3,W2W_1, W_3, W_2W1​,W3​,W2​),比原来多一个。为了保持参数量不变,中间层维度要缩小:

隐层维度=83d≈2.67d(保持参数量与 FFN(4d) 相当)\text{隐层维度} = \frac{8}{3}d \approx 2.67d \quad \text{(保持参数量与 FFN(4d) 相当)} 隐层维度=38​d≈2.67d(保持参数量与 FFN(4d) 相当)

LLaMA 用的就是这个:intermediate_size = 8192 而 hidden_size = 4096,比例是 2.0 而不是 4.0。

验证:

FFN 类型中间维度参数量(d=4096d=4096d=4096)
FFN(4d)163843×d23 \times d^23×d2
SwiGLU(8/3 d)109233×d×(8/3)d=8d23 \times d \times (8/3)d = 8d^23×d×(8/3)d=8d2
LLaMA81923×4096×8192=100M3 \times 4096 \times 8192 = 100M3×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 组★
1
2
3

LLaMA2 用 MQA,LLaMA3 用 GQA。 GQA 是 MHA 和 MQA 的折中,效果接近 MHA、成本接近 MQA。


七、LLaMA 完整配置对照表

组件原始 TransformerLLaMA / 现代 LLM原因
结构Encoder-DecoderDecoder-only统一任务,训练并行
归一化位置Post-LNPre-LN梯度稳定
归一化类型LayerNormRMSNorm可融合,更快
位置编码正弦绝对位置RoPE相对位置,外推好
FFNReLUSwiGLU门控表达力强
FFN 中间维度4d~2.67d补偿 SwiGLU 的额外矩阵
激活ReLUGELU / SiLU平滑
AttentionMHAGQA(推理友好)省 KV cache
初始化XavierN(0,0.02)N(0, 0.02)N(0,0.02)RMSNorm 已处理尺度
归一化无 bias无 bias更简洁
学习率调度warmup+decaywarmup + cosine标准配置
精度fp32bf16/fp16效率

八、从零实现一个 LLaMA Block

把上面所有内容串起来:

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

注意几个数字:

  1. 初始 loss = 10.50。理论上 ln⁡(32000)=10.37\ln(32000) = 10.37ln(32000)=10.37——我们的 loss 略高于均匀分布,因为用了 std=0.02 的初始化。训练开始时 loss 应该在 ln⁡(V)\ln(V)ln(V) 附近,这是初始化正确性的快速检查。

  2. 参数量 43.8M 而 4 层模型的主体(不含 embedding)只有约 1.1M。embedding 占了绝大部分(32000 × 512 = 16.4M,两个 embedding 32.8M)。这是小模型的典型特征。

  3. 代码里 GQA(n_kv_heads=2)能正常工作——这是实现中最容易出错的地方,我第一版忘了在 attention 之前把 KV 头repeat 到和 Q 头一样多,报形状不匹配的错误。

想验证 GQA 是否真的省了显存,对比一下:

python
# 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()):,} 参数")
1
2
3
4
5
6
7
8

实测输出:

MHA: 45,353,472 参数
GQA: 43,780,608 参数
省 1,572,864 参数
1
2
3

GQA 省了 157 万参数(占 transformer 部分的约 3.5%)。这个比例在真实模型里更明显——因为真实模型的 transformer 部分占比更高,且序列更长(KV cache 的节省与序列长度成正比)。


九、LLaMA2 相比 LLaMA 改进了什么

改进说明
更多数据2T tokens(LLaMA 是 1.3T)
上下文 4KLLaMA 是 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 → 下一层
1

最后一层的输出从未经过归一化(norm 在每个 block 的内部)。所以最后一个 block 的输出尺度是任意的——可能很大或很小,直接送进 lm_head 会导致 logits 尺度失控。

python
x = layer(x) for _ in range(N)
x = self.norm(x)         # ★ 必须
logits = self.lm_head(x)
1
2
3

Post-LN:不需要。

x → sublayer(x) + x → norm → 下一层
1

每个 block 的输出都经过了 LayerNorm,最后那个 block 的输出已经归一化过了。

python
x = layer(x) for _ in range(N)
logits = self.lm_head(x)   # 直接用,因为已经归一化
1
2

忘记加 final norm 是 Pre-LN 实现最常见的 bug——症状是训练 loss 曲线很奇怪,或者 logits 数值过大导致 softmax 完全饱和。

Q2:SwiGLU 有三个矩阵,为什么中间层维度反而要缩小?

答案

因为要保持参数量不变。

原始 FFN(两个矩阵):

参数量=d×4d+4d×d=8d2\text{参数量} = d \times 4d + 4d \times d = 8d^2 参数量=d×4d+4d×d=8d2

SwiGLU(三个矩阵),设中间维度为 hhh:

参数量=d×h+d×h+h×d=3dh\text{参数量} = d \times h + d \times h + h \times d = 3dh 参数量=d×h+d×h+h×d=3dh

令两者相等:

3dh=8d2⇒h=83d≈2.67d3dh = 8d^2 \quad \Rightarrow \quad h = \frac{8}{3}d \approx 2.67d 3dh=8d2⇒h=38​d≈2.67d

如果保持原来的 4d4d4d,参数量会变成 3×d×4d=12d23 \times d \times 4d = 12d^23×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 的优势:

  1. 训练可以完全并行——causal mask 让每个位置同时预测下一个词,整句话一次算完。encoder-decoder 也并行,但 decoder 只能看到 encoder 输出,任务设计更复杂
  2. 统一任务——加 [MASK] token 就变成 BERT 的填空任务,加 [BOS] 就变成生成任务。一个架构做所有事
  3. 架构简单——只有一种 block,不需要维护两套参数
  4. 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)N(0, 0.02)N(0,0.02) 初始化所有层,这不会有问题吗?

答案

不会,而且这正是架构设计的好处——因为 RMSNorm 已经把尺度问题解决了。

对比需要精细初始化的场景:

传统 MLP 中,每层的输出尺度会逐层累积。第 1 层输出 std 是 0.02,第 2 层就是 0.02 × (权重尺度)……如果不精心设计初始化,深层网络必然出问题(第 9 篇的实验:默认初始化衰减 4.65e14 倍)。

为什么 LLaMA 敢用统一初始化:

  1. RMSNorm 在每个子层入口把激活归一化了 → 进入下一个 Linear 时方差恒为 1 → 不需要考虑跨层的尺度累积
  2. Pre-LN 残差通路是干净恒等 → 梯度稳定
  3. 所有层形状相似(hidden_dim 统一)→ 一个 std 够用

0.02 这个数字的来源:

python
nn.init.normal_(m.weight, mean=0.0, std=0.02)
1

经验值,比 Xavier 略大一点。GPT-2 用的也是 0.02。实际差别不大,0.01~0.02 都能训。

这个案例的启示:好的架构设计(加归一化)能让初始化这个「魔法数字」变得不重要。 这比调参重要得多。


下一篇 → 推理优化:KV Cache 与 GQA

上一级: 目录 · 上一篇

13 · 推理优化:KV Cache 与 GQA

核心问题:自回归生成为什么慢?KV Cache 省的是什么?GQA 和 MQA 怎么权衡?


一、自回归生成的困境

1.1 问题的本质

LLM 生成文本是逐个 token 的:生成第 ttt 个 token 需要前 t−1t-1t−1 个作为输入。

朴素做法:每生成一个 token,就把整个序列重新前向一遍。

生成 nnn 个 token 的总计算量:

朴素=∑t=1nO(t2d)=O(n3d)\text{朴素} = \sum_{t=1}^{n} O(t^2 d) = O(n^3 d) 朴素=t=1∑n​O(t2d)=O(n3d)

这是一个 O(n3)O(n^3)O(n3) 的算法。 生成 1000 个 token,要做约 33 万倍于单步的计算——其中绝大部分被浪费了。

1.2 关键洞察:K 和 V 可以复用

观察注意力的输入:

是否依赖已生成的内容
QQQ(当前 token 的查询)✅ 每次都不同
KKK(已生成 token 的键)❌ 对已生成的 token,永不改变
VVV(已生成 token 的值)❌ 对已生成的 token,永不改变

这是因为 causal mask:token jjj 的 kj,vjk_j, v_jkj​,vj​ 只依赖它自己和它前面的内容,不依赖后面的 token。

所以生成第 ttt 个 token 时,k1…kt−1,v1…vt−1k_1 \dots k_{t-1}, v_1 \dots v_{t-1}k1​…kt−1​,v1​…vt−1​ 和上次生成第 t−1t-1t−1 个 token 时完全一样。

把它们存下来复用,就是 KV Cache。


二、KV Cache 的效果

2.1 计算量对比

无 cache:生成第 t 个 token 要算 t×t 个注意力分数
有 cache:生成第 t 个 token 只算 t×1 个注意力分数
1
2

生成 nnn 个 token 的总量:

方案总计算量n=1000n=1000n=1000 时的相对比
无 cacheO(n3)O(n^3)O(n3)10910^9109
有 cacheO(n2)O(n^2)O(n2)10610^6106

降了 1000 倍(nnn 倍)。

2.2 但注意力不是瓶颈:MLP 才是

一个容易误解的点:KV Cache 优化的是注意力部分,但推理的真正瓶颈是 MLP(FFN)。

LLaMA-7B 的参数量分布:

组件参数量占比
Attention (QKV + O)4d2=4×40962=67M4d^2 = 4 \times 4096^2 = 67M4d2=4×40962=67M9.5%
MLP (SwiGLU)3d×8192=101M3 d \times 8192 = 101M3d×8192=101M14.3%
其他(embedding + lm_head)~640M76%

每生成一个 token,所有权重都要参与计算(这是无法避免的),但 attention 的 KV 只算一次。

所以:

技术优化的部分实际收益
KV Cacheattention 的 K,VK,VK,V 投影中等(省了重复计算)
减少权重(量化/剪枝)所有权重大
算子融合显存访问大

这是为什么 4-bit 量化(GPTQ/AWQ)比 KV Cache 更能提升吞吐。

2.3 KV Cache 的代价:显存

KV Cache 的显存占用:

KV cache=2×nkv×L×dhead×nlayers×bytes\text{KV cache} = 2 \times n_{kv} \times L \times d_{head} \times n_{layers} \times \text{bytes} 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]    最省,质量下降
1
2
3

GQA 的思路:把 hhh 个 query 头分成 ggg 组,每组共享一个 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 个一组
1
2
3
4
5
6
7

3.2 GQA 实现

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

3.3 实测:GQA 的质量-成本权衡

python
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")
1
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.0x
1
2
3
4
5
6
7
8

3.4 各模型的 GQA 配置

模型n_headsn_kv_heads分组数KV cache 相对 MHA
LLaMA-1 7B32321(=MHA)100%
LLaMA2 70B648812.5%
LLaMA2 7B/13B32/4032/401(=MHA)100%
Llama3 8B328425%
Llama3 70B648812.5%
Mistral 7B328425%
Qwen2 7B284714.3%

注意一个现象:即使同一个家族,不同尺寸的模型 GQA 配置也不同(LLaMA2-70B 用 8 组,13B 用 MHA)。因为大模型推理时的显存压力更大,更需要省 KV cache。

实测研究数据(GQA 论文):hg=4∼8\frac{h}{g} = 4\sim 8gh​=4∼8 时,质量接近 MHA,节省接近 MQA。这是目前的最优权衡点。


四、Prefill vs Decode:两个阶段

现代推理框架把生成拆成两个阶段,优化完全不同:

4.1 Prefill(预填充)阶段

处理整个输入 prompt,一次前向算出所有 token 的 KV。

输入: [t1, t2, ..., t_2048]  →  一次前向  →  KV cache 填满
1

特点:

  • 计算密集(compute-bound)
  • 所有 attention 分数一次算完(可并行)
  • 占用 GPU 算力
  • 可以用 Flash Attention(第 14 篇)

瓶颈:GEMM 矩阵乘法,带宽利用率高但矩阵规模大。

4.2 Decode(解码)阶段

逐个生成 token。

cache: [k1..k2048]  +  新token  →  单步前向  →  cache 增长 1
1

特点:

  • 内存带宽受限(memory-bound)
  • 每步只算 1 个 query,矩阵规模极小(GEMM 形状是 [1,d]×[d,L][1, d] \times [d, L][1,d]×[d,L])
  • GPU 算力严重闲置——一个 32 核 GPU 大部分时间在等显存读数据
  • attention 部分计算量 O(L)O(L)O(L),但访存量 O(L)O(L)O(L)

瓶颈:显存带宽,不是算力。

4.3 为什么这个区分很重要

PrefillDecode
瓶颈算力显存带宽
关键优化Flash Attention、矩阵 kernel 优化量化、减小 KV cache
并行度高(所有位置并行)低(逐个)
优化手段Flash Attention、PagedAttention量化、GQA、投机解码

这解释了几个现象:

  1. 为什么量化(4-bit)对长文本生成帮助大——Decode 阶段是带宽瓶颈,权重量化直接减少要读的字节数
  2. 为什么 Continuous Batching 有效——把多个请求的 Decode 步骤拼成一个 batch,提高 GPU 利用率
  3. 为什么投机解码(Speculative Decoding)有效——用小模型并行猜多个 token,大模型一次验证,把带宽受限的串行变成算力更密集的批量

4.4 实测:Decode 阶段的带宽瓶颈

python
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")
1
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/s
1
2
3
4
5

计算依据:每层 4d2+3d×8192=67M+101M=168M4d^2 + 3d \times 8192 = 67\text{M} + 101\text{M} = 168\text{M}4d2+3d×8192=67M+101M=168M,32 层共 5.37B5.37\text{B}5.37B(不含 embedding 和 lm_head)。

结论:

精度权重占用A100 带宽上限
fp1610.7 GB186 tok/s
int85.4 GB373 tok/s
int42.7 GB745 tok/s

int4 相对 fp16 正好快 4 倍——纯粹来自「读的字节数变成 1/4」。

而 A100 的算力是 312 TFLOPS,远高于 186 tok/s 这个数字(差三个数量级)。这就是「Decode 是带宽瓶颈」的定量证明。


五、其他推理优化技术

技术原理收益
KV Cache复用已算的 K/V计算量 O(n3)→O(n2)O(n^3) \to O(n^2)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 jjj,它的 kj,vjk_j, v_jkj​,vj​ 的计算过程是:

kj=WKxj+RoPE(xj,j),vj=WVxjk_j = W_K x_j + \text{RoPE}(x_j, j), \qquad v_j = W_V x_j kj​=WK​xj​+RoPE(xj​,j),vj​=WV​xj​

只依赖 xjx_jxj​ 自己(以及它的位置),不依赖任何后续 token。

  • 生成 token ttt 时,k1…kt−1k_1 \dots k_{t-1}k1​…kt−1​ 的计算结果和生成 token t−1t-1t−1 时完全相同 → 可以缓存
  • 而 qtq_tqt​ 每次都是新计算的(新 token 的新 query)

反例(如果没有 causal mask):如果 attention 是双向的,token 1 的 k1k_1k1​ 也会依赖 token 2 的内容——但 token 2 还不存在,无法预先计算。这就是 KV Cache 依赖自回归结构的原因。

一个推论:encoder 模型(双向 attention)无法用 KV Cache,因为每层的 K/V 都依赖整个输入,无法增量计算。

Q2:KV Cache 让复杂度从 O(n3)O(n^3)O(n3) 降到 O(n2)O(n^2)O(n2),但为什么实际推理还那么慢?

答案

因为注意力从来不是瓶颈。

分析 LLaMA-7B 每生成一个 token 的计算量:

操作FLOPs占比
QKV 投影3×2d×L3 \times 2d \times L3×2d×L小
Attention(含 KV cache)4×d×L4 \times d \times L4×d×L(不是 L2L^2L2)很小
MLP (SwiGLU)3×2d×8192=6d2≈2×1083 \times 2d \times 8192 = 6d^2 \approx 2 \times 10^83×2d×8192=6d2≈2×108~70%
lm_head2×d×32000≈2.6×1082 \times d \times 32000 \approx 2.6 \times 10^82×d×32000≈2.6×108~28%

注意:

  • 有 KV cache 时,注意力是 O(L)O(L)O(L) 而不是 O(L2)O(L^2)O(L2)——原本的 L2L^2L2 项已经省掉了
  • MLP 的参数量是固定的,每个 token 都必须完整算一遍(不能缓存,因为没有重复利用)
  • lm_head 同理

真正的事实:Decode 阶段是显存带宽瓶颈,不是算力瓶颈。

每步要读 5.37B×2=10.75.37\text{B} \times 2 = 10.75.37B×2=10.7 GB 权重(fp16)。A100 的带宽约 2 TB/s,所以:

理论上限=2×101210.7×109≈186 token/s\text{理论上限} = \frac{2 \times 10^{12}}{10.7 \times 10^9} \approx 186 \text{ token/s} 理论上限=10.7×1092×1012​≈186 token/s

而 A100 的算力是 312 TFLOPS,差了三个数量级。GPU 大部分时间在等显存。

结论:想提速就该 减少要读的字节数(量化)或 增加并行度(continuous batching),而不是优化注意力。

Q3:GQA 的分组数 ggg 怎么选?为什么不是越大越好?

答案

ggg 的范围:

ggg等价于KV cache质量
hhh(=32)MHA100%基准
8GQA25%≈ MHA
4GQA12.5%略降
1MQA3%明显下降

论文(GQA, Ainslie et al. 2023)的结论:hg=4∼8\frac{h}{g} = 4 \sim 8gh​=4∼8 时质量接近 MHA。

为什么不能 g=1g=1g=1(MQA):

  1. 质量受损——实测困惑度上升 0.1~0.3,且在某些任务上明显
  2. 表达能力受限——所有 query 头共享一个 KV,意味着它们的「信息来源」完全相同。多头之所以有效,是因为每个头能看到不同的信息(第 11 篇)

为什么不能 g=hg=hg=h(MHA):

没有节省任何东西。KV cache 占用是最大瓶颈之一。

选择依据:

场景建议 ggg
短上下文(<2K)、追求质量g=hg = hg=h(MHA)
中等上下文(4-8K)h/g=4h/g = 4h/g=4
长上下文(32K+)、高并发h/g=8h/g = 8h/g=8
模型规模大倾向更小的 ggg(显存压力)

实际观察:LLaMA2-70B(64 头)用 g=8g=8g=8,而 LLaMA2-13B(40 头)用 MHA。大模型因为显存压力大,被迫用更激进的 GQA。

Q4:Prefill 和 Decode 阶段的优化手段完全不同,为什么?

答案

因为两个阶段的瓶颈完全不同——这是硬件特性决定的。

Prefill(处理整个 prompt)

形状:[batch, seq, d] × [d, d] 的矩阵乘法

GEMM: [2048, 4096] × [4096, 4096]
1

特点:

  • 矩阵很大,算术强度高(每读 1 字节能做很多次运算)
  • GPU 算力被充分利用
  • 瓶颈 = 算力

优化手段:

  • Flash Attention(减少显存读写,同时加速)
  • 更高效的 GEMM kernel
  • Tensor Core(fp16/bf16)

Decode(逐个生成)

形状:[1, 4096] × [4096, 4096](只有 1 个 token!)

GEMM: [1, 4096] × [4096, 4096]    ← 极度瘦长
1

特点:

  • 矩阵很瘦,算术强度极低——每读 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)≈(NcN)αN,αN≈0.076L(N) \approx \left(\frac{N_c}{N}\right)^{\alpha_N}, \quad \alpha_N \approx 0.076 L(N)≈(NNc​​)αN​,αN​≈0.076

其中 NNN 是参数量,NcN_cNc​ 是「消除损失所需的参数量」,αN\alpha_NαN​ 是幂律指数。

Hoffmann 定律(Chinchilla,2022):

给定算力 CCC,最优的参数量和训练 token 数满足:

Nopt∝C0.55,Dopt∝C0.45\boxed{N_{opt} \propto C^{0.55}, \qquad D_{opt} \propto C^{0.45}} Nopt​∝C0.55,Dopt​∝C0.45​

含义:算力翻倍时,参数量涨 20.55≈1.462^{0.55} \approx 1.4620.55≈1.46 倍,数据量涨 20.45≈1.372^{0.45} \approx 1.3720.45≈1.37 倍。

1.2 Chinchilla 定律的重大修正

之前的共识(GPT-3 时代):「参数比数据重要,所以应该把所有算力投在放大模型上」。

Hoffmann 的实验推翻了这个直觉。他训练了 50 个模型(从 400 万到 660 亿参数),发现:

2021 年的 OPT-175B 是「严重欠训练」的——按 Chinchilla 定律,它应该用 4 倍的数据量训练。

模型参数量Token 数Token/参数是否最优
GPT-3175B300B1.7❌ 欠训练
Chinchilla70B1.4T20✅ 最优
Llama 2 7B7B2T286✅ 远超最优
Llama 3 8B8B15T1875✅★ 远超

关键观察:2023 年后的模型(Llama 系列)Token/参数比远高于 Chinchilla 最优值。

这是不是因为违反定律? 不是。原因是推理成本:

  • Chinchilla 定律假设「训练算力和推理算力同等重要」
  • 但实际部署中,一个模型要被推理几百万次。参数量直接决定推理成本(O(N)O(N)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 的计算量
1
2
3
4
5
6
7
8
9
10

注意:小数据集上正确的做法是训小模型 + 多epoch,而不是训大模型(会严重过拟合)。

1.4 缩放定律的局限(重要)

不要盲目迷信缩放定律:

  1. 它只描述「趋势」,不预测具体数值。知道「10B 模型在 100B token 下 loss 大约是 X」很难实际用到。
  2. 架构和数据质量的差异会被掩盖。一个架构更好的模型可能比参数多 2 倍的模型更好。
  3. 存在「能力涌现」现象(emergent abilities)——某些能力在规模到一定程度前几乎不出现。Wei et al. (2022) 发现 100B 参数以下模型在某些推理任务上接近随机猜测。

关于能力涌现的一个重要修正:Schaeffer et al. (2023) 指出,很多「涌现」是评测指标的选择问题——用连续指标(如准确率)测量会显示突变,但用「与随机猜测的差距」这类归一化指标测量,会显示平滑的增长曲线。

实践意义:不要因为「我的小模型没有涌现能力」就认为方法有问题。 可能只是没到那个规模,或者指标选错了。


二、高效注意力:五条技术路线

2.1 问题回顾

标准注意力有 O(L2)O(L^2)O(L2) 的时间和空间复杂度(第 10 篇)。

完整注意力矩阵: [B, H, L, L]
LLaMA-7B, B=2, H=32, L=4096, fp16 → 2.15 GB/层 → 32层 = 68 GB
1
2

两条完全不同的思路:

思路代表做法
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]                  ← 慢
1
2
3
4

关键数据:A100 的 HBM 带宽约 2 TB/s,但 SRAM(片上缓存)带宽约 19 TB/s——快 10 倍。

FlashAttention 的想法:把计算分成小块,让中间结果尽量待在 SRAM 里,只把最终需要的写回 HBM。

3.2 分块计算(Tiling)

把 Q,K,VQ, K, VQ,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  # 累加,不写回HBM
1
2
3
4
5
6
7
8

关键点:softmax\text{softmax}softmax 的分母(所有行的指数和)可以增量累加:

ℓinew=max⁡(ℓiold,max⁡jSij),minew=ℓineweℓiold−ℓinewmiold+∑jeSij−ℓinew\ell_i^{new} = \max(\ell_i^{old}, \max_j S_{ij}), \quad m_i^{new} = \ell_i^{new} e^{\ell_i^{old} - \ell_i^{new}}m_i^{old} + \sum_j e^{S_{ij} - \ell_i^{new}} ℓinew​=max(ℓiold​,jmax​Sij​),minew​=ℓinew​eℓiold​−ℓinew​miold​+j∑​eSij​−ℓinew​

这就是 online softmax 的技巧——不用等所有 KKK 块都算完就能得到部分结果。

3.3 反向传播的重计算

一个疑问:前向不存注意力矩阵了,反向怎么办?

答案:重计算(recomputation)。反向传播时:

  1. 只存了 OOO(输出)、LLL(每行的 logsumexp)、MMM(每行的 max)
  2. 反向时从 HBM 重新读 Q, K, V,重新分块算
  3. 用 L,ML, ML,M 恢复 softmax 分母

代价:反向时多算一遍 QKTQK^TQKT(约增加 30% FLOPs),换来显存降几十倍。

3.4 实测:FlashAttention 的收益

python
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")
1
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            409x
1
2
3
4
5
6
7

关键观察:序列越长,节省倍数越大——因为完整矩阵是平方增长,FlashAttention 是线性。

实际收益总结(LLaMA-7B,2048 token):

指标标准实现FlashAttention
显存100%~30%
速度100%140~200%
计算结果精确完全精确

注意「速度 140~200%」——FlashAttention 不只省显存,还更快,因为它减少了 HBM 读写。

3.5 用法

python
import torch.nn.functional as F

# ★ 一行代码,PyTorch 会自动选最优后端
out = F.scaled_dot_product_attention(q, k, v, is_causal=True)

# 训练时建议也用它(支持 backward)
1
2
3
4
5
6

实测速度(如果你的 GPU 支持):

python
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 才能测出加速比")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26

四、近似注意力:五条路线

这些方法真的改变了数学,用近似换复杂度。

4.1 Linformer:低秩投影

Attention(Q,K,V)≈softmax(QK⊤EF)(EGV)\text{Attention}(Q,K,V) \approx \text{softmax}(Q K^\top E_F)(E_G V) Attention(Q,K,V)≈softmax(QK⊤EF​)(EG​V)

其中 EF,EGE_F, E_GEF​,EG​ 把 K,VK, VK,V 从 LLL 维投影到 kkk 维(k≪Lk \ll Lk≪L)。

复杂度:O(Lkd)O(Lk d)O(Lkd) —— 线性。

致命缺陷:只能用于固定的序列长度。EF∈Rk×LE_F \in \mathbb{R}^{k \times L}EF​∈Rk×L 依赖 LLL,训练在 512 上就不能用于 1024。

4.2 Performer:线性注意力

核心技巧:改变矩阵乘法的顺序。

softmax(QK⊤)V≈ϕ(Q)(ϕ(K)⊤V)\text{softmax}(QK^\top)V \approx \phi(Q)(\phi(K)^\top V) softmax(QK⊤)V≈ϕ(Q)(ϕ(K)⊤V)

因为 ϕ\phiϕ 可以分离,softmax(a+b)=softmax(a)softmax(b)\text{softmax}(a+b) = \text{softmax}(a)\text{softmax}(b)softmax(a+b)=softmax(a)softmax(b) 近似成立。

复杂度:先算 ϕ(K)⊤V\phi(K)^\top Vϕ(K)⊤V(k×Vk \times Vk×V,与 LLL 无关),再左乘 ϕ(Q)\phi(Q)ϕ(Q)。

但这里有个关键难点——softmax 不能这样分离:

softmax(QK⊤)≠ϕ(Q)ϕ(K)⊤\text{softmax}(QK^\top) \ne \phi(Q)\phi(K)^\top softmax(QK⊤)=ϕ(Q)ϕ(K)⊤

Performer 的解法:用 FAVOR+ 核。

ϕ(x)=elu(x)+1\phi(x) = \text{elu}(x) + 1 ϕ(x)=elu(x)+1

这个正核(positive kernel)的设计使得 softmax 的指数核可以被无界地近似:

exp⁡(q⊤k)≈ϕ(q)⊤ϕ(k)\exp(q^\top k) \approx \phi(q)^\top \phi(k) exp(q⊤k)≈ϕ(q)⊤ϕ(k)

用无偏随机估计做期望的蒙特卡洛估计。

致命缺陷:精度损失明显。实测困惑度比标准注意力差,尤其在长序列上。已被主流放弃。

4.3 Longformer / BigBird:稀疏模式

只算部分位置,保证每个 token 能看到足够多的上下文:

滑动窗口(局部)    ●●●●●●○○○●●●●●●○○○●●●●●●
全局注意力        ●●●●●●●●●●●●●●●●●●●●●●●●●●●●●
                   ↑↑↑   ↑
                局部窗口  少数全局 token
1
2
3
4

复杂度:O(Lw)O(Lw)O(Lw),www 是窗口大小。

优势:精度损失小,可以做到几乎无损。

劣势:必须固定窗口模式,灵活性差。已经被 FlashAttention 取代(因为 FlashAttention 既精确又省)。

4.4 线性注意力的一般视角

所有线性注意力都遵循同一个模式:

softmax(QK⊤)V→ϕ(Q)(ϕ(K)⊤V)⏟只算一次,与 L 无关\text{softmax}(QK^\top)V \to \phi(Q)\underbrace{(\phi(K)^\top V)}_{\text{只算一次,与 } L \text{ 无关}} softmax(QK⊤)V→ϕ(Q)只算一次,与 L 无关(ϕ(K)⊤V)​​

改变计算顺序:先算 K⊤VK^\top VK⊤V(维度 d×dd \times dd×d,很小),再算 Q(K⊤V)Q(K^\top V)Q(K⊤V)。

方法ϕ\phiϕ精度
Linear Transformerelu + 1差
PerformerFAVOR+ 正核较差
RetNet分组归一化 +门控好(专为训练设计)
Mamba选择性状态空间好

Mamba(2023)是这个方向最成功的:用选择性状态空间模型(Selective SSM) 替代注意力,实现 O(L)O(L)O(L) 的序列建模,且不需要训练时的并行(可以像 RNN 一样流式推理)。

4.5 现状判断

方法是否主流原因
FlashAttention✅ 事实标准精确 + 更快 + 更省
PagedAttention (vLLM)✅ 主流(推理侧)解决 KV cache 碎片
Mamba / SSM🔶 特定场景适合超长序列、流式
Linear Attention🔶 部分特定场景
Performer / Linformer❌ 已淘汰精度损失明显
Longformer❌ 已淘汰被 FlashAttention 取代

核心判断:近似注意力的时代结束了,FlashAttention 赢了。 因为它证明了「精确计算 + IO 优化」的收益大于「近似计算 + 低复杂度」。

长序列的战场已经转移到:

  1. 稀疏 + 优化(FlashAttention + 稀疏模式,如 LongFlashAttention)
  2. 状态空间模型(Mamba 系列,完全不同的架构)
  3. 线性注意力 + 混合(Hybrid,如 Jamba、Zamba)

五、当前的架构演化方向

5.1 三条主线

方向代表核心思想
更长的上下文LongFlashAttention、RingAttention在 FlashAttention 基础上做序列并行
更省的状态Mamba / SSM用递推状态替代 KV cache
混合架构Jamba、Zamba、Qwen3部分层用 SSM,大部分层用注意力

5.2 值得关注的判断

目前没有「注意力将被替代」的证据。 主流仍是 Transformer + FlashAttention。SSM 是有前景的补充而非替代——因为:

  1. Attention 有 O(1)O(1)O(1) 的「随机访问」能力(可以关注任意位置),SSM 是递推的(只能看历史压缩)
  2. Transformer 的硬件生态(Tensor Core、FlashAttention kernel)成熟且高度优化
  3. 归纳偏置不同:SSM 假设序列有强局部结构(像 RNN),而语言有长距离依赖——attention 更适合

一个务实的判断:如果你现在要做 LLM 项目,Transformer + FlashAttention 仍是最安全的选择。


六、动手实验

实验 1:FlashAttention 的正确性验证

python
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 是精确算法,不是近似。")
1
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 是精确算法,不是近似。
1
2
3

实验 2:验证「近似注意力的精度损失」

python
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}倍    ← ★ 输出量级完全不同了")
1
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倍    ← ★ 输出量级完全不同了
1
2
3

这个实验的结果非常惊人:Linear Attention 的输出量级是标准注意力的 5000 多倍,相对误差达到 1.3 万倍。

根本原因:extsoftmax ext{softmax}extsoftmax 会把每一行归一化到和为 1(输出被约束在 [−1,1][-1, 1][−1,1] 量级),而 extelu(x)+1 ext{elu}(x)+1extelu(x)+1 可以任意大。两者的「值域」完全不同,直接换过去必然崩溃。

这就是为什么必须用正核(Performer 的做法)——它让 ϕ\phiϕ 的值域和 exp⁡(⋅)\exp(\cdot)exp(⋅) 匹配。但即使如此,误差仍然显著。

这就是 Performer 最终没被采用的原因:近似带来的误差超过了复杂度收益。

实验 3:缩放定律的可视化

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

关键观察:

  1. 损失下降非常平缓:参数量从 1M 涨到 100B(10万倍),loss 只从 2.09 降到 0.87。这就是为什么扩容是「暴力」但有效的方式

  2. 每增加 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. 写输出
1
2
3
4
5

至少 3 次完整的 O(L2)O(L^2)O(L2) 读写。

FlashAttention 的做法:分块计算,中间结果留在 SRAM:

for Qi:
    for Kj:
        在 SRAM 里算 Qi@Kj.T(不写 HBM)
        在线更新 softmax 累加器(online softmax)
        累加到 O[i]
    只写最终的 O[i] 到 HBM
1
2
3
4
5
6

HBM 访问从 O(L2)O(L^2)O(L2) 降到 O(L)O(L)O(L)。

代价:反向传播时需要重计算(多算一遍 QKTQK^TQKT,约 +30% FLOPs)。但用 30% 的额外计算换 10 倍的 IO 效率,净收益是正的。

实测效果:A100 上 FlashAttention-2 比标准实现快 140%~200%,同时显存从 100% 降到约 30%。

Q2:线性注意力为什么精度损失这么大?有没有根本的改进方向?

答案

根本障碍:softmax 不可分离。

Attention=softmax(QK⊤d)V\text{Attention} = \text{softmax}\left(\frac{QK^\top}{\sqrt d}\right)V Attention=softmax(d​QK⊤​)V

线性注意力想做的:

softmax(QK⊤)V≈?ϕ(Q)(ϕ(K)⊤V)\text{softmax}(QK^\top)V \stackrel{?}{\approx} \phi(Q)(\phi(K)^\top V) softmax(QK⊤)V≈?ϕ(Q)(ϕ(K)⊤V)

这要求 softmax(ab⊤)=ϕ(a)ϕ(b)⊤\text{softmax}(ab^\top) = \phi(a)\phi(b)^\topsoftmax(ab⊤)=ϕ(a)ϕ(b)⊤——但 softmax 对每一行独立归一化,无法写成两个向量的外积。

这个障碍是本质的:归一化操作打破了矩阵分解结构。

改进方向:

  1. Performer 的正核(FAVOR+):用 ϕ(x)=elu(x)+1\phi(x) = \text{elu}(x) + 1ϕ(x)=elu(x)+1 这样的正核,让 exp⁡(q⊤k)≈ϕ(q)⊤ϕ(k)\exp(q^\top k) \approx \phi(q)^\top\phi(k)exp(q⊤k)≈ϕ(q)⊤ϕ(k) 成立。误差大幅减小,但仍不如精确

  2. RetNet 的分组归一化:不用 softmax,改用可分离的归一化(分组 + 组内归一化)。训练效果接近标准注意力,且支持并行训练

  3. FlashAttention 的正交路径:不改数学,只优化实现。这是当前最优解

结论:只要硬件和IO技术还能优化,「精确计算」就永远优于「近似计算」。 近似算法是在硬件能力不足时的妥协,现在不再是必要。

判断标准:如果你的序列长度 < 32K,用 FlashAttention 就够了,不要碰近似算法。

Q3:缩放定律说「算力最优分配是 N∝C^0.55, D∝C^0.45」,但为什么 Llama 3 的 Token/参数比远超这个最优值(1875 vs 20)?

答案

因为 Chinchilla 定律的假设和实际场景不符。

Chinchilla 定律的隐含假设:训练成本和推理成本的权重相同。

C=6ND(训练 FLOPs)C = 6ND \quad \text{(训练 FLOPs)} C=6ND(训练 FLOPs)

实际部署的权衡:

训练一次→ 模型被推理几百万次

推理成本 ∝ 参数量 N(每次调用)
训练成本 ∝ N × D(一次)

所以总成本 = N × D + N × (推理次数)
1
2
3
4
5
6

当推理次数是百万量级时,NNN 的权重远超 DDD。最优策略应该偏向更大的 NNN、更小的 DDD。

Llama 3 的具体做法:

模型参数量Token 数Token/参数Chinchilla 最优
Llama 2 7B7B2T28620(超14倍)
Llama 3 8B8B15T187520(超94倍)

这是有意的「过训练」——牺牲训练算力换取推理效率。

换算成实际成本:训练 Llama 3 8B 需要约 15T token,按 Chinchilla 最优只需要 160B token。多花了 94 倍的算力。但因为模型只有 8B,推理便宜得多——在 billions of requests 的场景下,这笔交易极其划算。

这个案例的教学意义:缩放定律是描述「训练侧」最优的工具,而实际决策要考虑训练和推理的全生命周期成本。 机械套用会做出错误决策。

Q4:Mamba 这类 SSM 模型会取代 Transformer 吗?

答案

目前不会,但值得持续关注。

SSM 的优势:

  1. O(L)O(L)O(L) 复杂度,长序列上比 attention 快得多(O(L2)O(L^2)O(L2))
  2. 推理状态是常数大小(Mamba-2 的状态是固定大小),不需要 KV cache
  3. 流式推理天然支持——RNN 式,逐 token 处理,显存不随长度增长

SSM 的劣势:

  1. 训练需要串行(RNN 式),不能像 attention 那样完全并行。这是致命伤——GPU 喜欢并行
  2. 归纳偏置受限——SSM 是递推的,只能压缩历史信息到固定大小的状态。而 attention 可以「随机访问」任意位置
  3. 长距离依赖处理能力弱——实证显示 SSM 在需要精确回忆的任务上不如 attention

为什么现在的答案是「混合」:

  • Jamba:大部分层是 Transformer,每隔几层插一个 Mamba 层
  • Zamba:Mamba + 共享的 Transformer 层并行
  • Qwen3:不同规模用不同架构

这是最务实的方向:取两者之长——attention 负责精确的长距离依赖,SSM 负责高效的局部/层次处理。

我的判断(如果要我下注):

  1. 未来 2-3 年:Transformer + FlashAttention 仍是主流,SSM 只在特定场景(超长序列、流式、边缘设备)占优
  2. 更长期的变数:如果出现某种硬件/算法突破让串行训练不昂贵,SSM 才有机会
  3. 对你的实际决策:如果现在做 LLM 项目,不要为了「未来可能」而用 SSM。Transformer 的生态、工具、预训练权重、量化方案都是压倒性优势

判断新架构是否值得跟进的三个问题:

  1. 有没有开源的预训练权重?(没有 = 你要从零训 = 基本不现实)
  2. 推理成本真的更低吗?(要看 Prefill 和 Decode 分开算)
  3. 在你的具体任务上验证过吗?(不要信 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 分钟,判断要不要继续读

目标:搞清这篇论文「声称」做了什么。

只读:

  1. 标题 + 摘要(Abstract)
  2. ** introduction**(通常 1-2 页)
  3. 看所有图表(Figure 1~3,Table 1)
  4. 结论(Conclusion)

这一遍要能回答:

  • 解决什么问题?
  • 声称的方法是什么?(一句话)
  • 声称的效果如何?
  • 有没有代码/数据链接?

决策点:如果这三遍之后你觉得「这和我在做的事没关系」,就停下。论文不是书,不需要读完。

⚠️ 关键纪律:第一遍不许读方法和实验的细节。 直接翻到实验部分是最大的时间浪费。

第二遍:30-60 分钟,搞懂「怎么做的」

这一遍才开始读正文。

重点读:

  1. 方法章节的核心公式和图
  2. 实验设置的表格(尤其消融实验)
  3. 论文里提到的所有 baseline

要能回答:

  • 方法的核心思想是什么?
  • 它和最相关的已有工作有什么本质区别?
  • 关键的设计选择是什么?为什么这样选而不是那样?
  • 消融实验说明了哪些部分是必要的?

⚠️ 重点:带着问题读。每次读之前先问自己「我想搞懂什么」,然后在论文里找答案。

第三遍:2-4 小时,验证「真的有用吗」

这一遍才是真正的检验。

要做的事:

  1. 重现论文的实验(至少跑通一个 baseline)
  2. 找出论文没说的限制
  3. 想清楚这个方法在什么情况下会失效
  4. 判断能不能用在你的问题上

要能回答:

  • 报告的结果可信吗?(数据划分公平吗?超参是对比方法调好的吗?)
  • 消融实验支持作者的结论吗?
  • 这个方法的代价是什么?(参数量、显存、训练时间——论文往往不说)
  • 我能用它做什么?

三、论文的结构解剖(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                              ← 需要时再读(超参、实现细节)
1
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 打不过
只在自己的数据集上测说明泛化性存疑
只报最好的一次 runcherry-picking
「我们达到了 SOTA」但只在一个数据集泛化性存疑
消融实验里某模块去掉性能更好那个模块是负担
方法描述模糊,无法复现可能只是调参技巧

4.3 「可疑但不是造假」的情况

这些是常见问题,不一定是故意的,但你要知道:

问题 1:对比超参不公平

这是论文领域最普遍的问题。不是造假,但结论不可靠。

应对:自己复现时,给每个方法都认真调参,不要用作者报告的数字。

问题 2:消融实验不完整

作者只消融了「显眼」的模块,没消融那些微妙的实现细节(学习率、初始化、数据增强)。

应对:把消融实验当作「论文声称的证据」,不是「完整的证明」。

问题 3:只在特定超参下有效

新方法可能只是「换了一种隐式正则化」,在作者的超参下表现好,换个超参就不如 baseline。

应对:至少试 2-3 组超参,看效果是否稳定。

问题 4:数据泄漏

预处理时用了全量数据算统计量(标准化、词表构建、去重),导致测试集信息泄漏到训练。

应对:检查数据处理 pipeline(第 2 篇讨论过的泄漏类型)。


五、按需的资料检索

5.1 怎么找论文

渠道特点
Google Scholar引用数、被引用列表(找后续工作)
arXiv最新,但未经 peer review
Semantic ScholarAI 辅助摘要,读不懂时有用
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" 复现  找复现
1
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-3
1
2
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 decoding
1
2
3
4
5
6

七、动手实践

练习 1:用三遍法读一篇论文

选一篇:ViT(An Image is Worth 16x16 Words),arXiv:2010.11929。

计时:

遍时间只看产出
第一遍8 分钟标题、摘要、引言、所有图、结论一句话总结 + 要不要继续
第二遍45 分钟方法全文、实验表格方法核心 + 消融结论
第三遍90 分钟跑代码(如果能)、找限制我的判断:能用在哪

记录(这是练习的核心产出):

markdown
## ViT (2020)
- 问题: 卷积网络有很强的归纳偏置,但卷积不适合大规模数据
- 方法: 把图像切成 patch,展平成序列,用标准 Transformer
- 关键发现: 需要更大的模型 + 大量数据才能超过 CNN
- 消融: patch size 越小越好;位置编码可有可无(但用了更好)
- 局限: 小数据集上不如 CNN;自注意力 O(L²)
- 我的判断: 证明了「Transformer 可以做视觉」,直接催生了后续的 CLIP/GPT-V
- 可迁移的点: 「归纳偏置 vs 数据规模」的权衡适用于所有领域
1
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%,你该怎么做才能确认这个改进是真的?

答案

五步验证:

  1. 算置信区间:测试集多大?1000 个样本时 1% 差异可能只是 10 个样本的差别

    SE=p(1−p)n\text{SE} = \sqrt{\frac{p(1-p)}{n}} SE=np(1−p)​​

    n=1000,p=0.85n=1000, p=0.85n=1000,p=0.85 时 SE≈1.1%\text{SE} \approx 1.1\%SE≈1.1% —— 1% 的差异完全在噪声范围内

  2. 跑多seed:至少 3 个随机种子,报均值 ± std。如果 std 有 2%,那 1% 的差异毫无意义

  3. 配对检验:两个模型在同一批测试样本上评估,用 McNemar 检验(分类任务)而不是独立样本 t 检验

  4. 给 baseline 也认真调参:如果 baseline 用的是论文里的默认超参而新方法调了很久,这个对比无效

  5. 看是否在其他任务/数据集上还有效:如果只在特定设置下有效,说明是脆弱的改进

如果做完全部五步还站得住,才能说这个改进是真实的。

实践中的简化:至少做第 2 和第 4 步。因为这两步成本最低而收益最高。

Q2:一篇论文的消融实验显示「去掉模块 C 后性能提升了 0.5%」,作者在正文里说「模块 C 有效」。这个说法有问题吗?

答案

有明显问题。作者的说法和自己的数据矛盾。

去掉一个模块性能反而更好,说明这个模块是有害的——至少在当前设置下。作者应该报告的是「去掉 C 更好」这个发现,而不是硬说 C 有效。

这种时候要警惕几种可能:

  1. 消融实验的实现有 bug(比如忘记同步改其他地方)
  2. 超参没重新调——去掉 C 后模型的最好超参变了,用原来的超参测不公平
  3. 0.5% 在噪声范围内——需要多 seed 验证
  4. 作者的叙述和实验脱节——论文写完才改数据,或者根本是笔误

正确做法:自己跑一遍「去掉 C」的实验,用相同的数据划分和调优方式。如果确实更好,那论文的这部分结论是错的。

这个例子说明了一件事:消融实验是论文里最可信的部分(因为做起来麻烦、不容易造假),所以它和正文叙述冲突时,相信消融实验。

Q3:怎么判断一个 arXiv 预印本的质量?等 peer review 是不是必须?

答案

不必等。 顶会(NeurIPS/ICML/ICLR)的论文从 arXiv 到会议决议通常有 6-12 个月,领域发展很快,等不起。

判断预印本质量的实用方法:

信号(好的):

  1. 作者背景:这个领域的活跃研究者,之前有靠谱的工作
  2. 有开源代码,且代码质量好(有 README、有复现说明、有 issue 讨论)
  3. 有社区讨论:HuggingFace / Reddit / Twitter 上有人在讨论甚至复现
  4. 后续被引用/被类似工作采用:如果 2024 年的论文引用了它,说明影响开始了
  5. 消融实验完整,局限讨论诚实

信号(差的):

  1. 只有 arXiv,投稿了但被拒
  2. 没有代码
  3. 消融实验缺失或明显敷衍
  4. 只在一个数据集上验证
  5. 声称 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 的粗搜)
1

成本很高(kkk 个配置 × mmm 个超参配置 × 训练成本),所以要取舍:

  • 优先级最高的配置(最核心的组件):仔细调
  • 次要配置:至少试 2-3 个 lr

一个有价值的中间做法:对所有配置用同一个 lr 搜索空间,都跑一遍,取各自最好的。这样既公平又成本可控。

2.4 交互作用:为什么需要组合消融

组件之间可能不是独立的。假设:

配置性能
Full90.0
− A89.5(掉 0.5)
− B88.0(掉 2.0)

表面上 B 比 A 重要。但如果 A 和 B 是配合使用的(比如 A 是 B 的前置条件),单独去掉任何一个都会破坏整个流程。

所以必须测组合:

配置性能
− A − B82.0(掉 8.0!)

这时候结论完全不同:A 和 B 各自看都不重要,但去掉两个就崩了。这说明它们是耦合的,必须一起用。

这是消融实验里最容易漏掉的情况,也是最能出洞见的地方。


三、如何验证一个「改进」

3.1 五步验证法

假设你实现了一个改进,想验证它真的有效:

Step 1:算统计显著性

python
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对
    ...
1
2
3
4
5
6
7
8
9

关键:用配对检验(paired test),因为两个模型在同一批样本上评估,样本间的差异可以消掉。

不要用独立样本 t 检验——那会低估显著性。

Step 2:多个随机种子

python
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}")
1
2
3
4

至少 3 个 seed,推荐 5 个。如果 std > 改进幅度,这个改进不可信。

Step 3:给 baseline 调同样的超参

这是最重要也最容易被忽略的一步。如果你的改进只是「换了一种隐式正则化」,在精心调参下 baseline 可能更好。

Step 4:换不同的数据集/任务验证

一个方法只在一个设置下有效,通常说明它是脆弱的。

至少在 2-3 个设置下验证。如果预算紧,至少要「一个主数据集 + 一个不同领域的数据集」。

Step 5:报告代价

改进不是免费的。要报告:

指标为什么重要
参数量影响存储和部分推理成本
FLOPs影响推理算力
实际推理延迟★ 最接近用户体验的指标
训练成本影响可复现性
显存占用影响可用性

很多论文只报告精度不提代价,这是重要的信息缺失。

3.2 一个常见错误:把「训练更好」当成「模型更好」

python
# ❌ 错误:只看训练 loss
for epoch in range(100):
    train_loss = train(model)
    print(f"epoch {epoch}: train loss {train_loss:.4f}")
1
2
3
4

问题:训练 loss 降得更低可能意味着过拟合更强。

正确:同时监控验证指标:

python
# ✅ 正确
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()
1
2
3
4
5
6
7
8

这个区分很重要:第 2 篇的 ResNet 权重衰减实验里,「训练误差更高但测试误差更好」是常态。训练误差不是好模型的判据。


四、避免数据泄漏

这是实验设计里最隐蔽也最致命的错误。 详细清单见第 2 篇,这里补充实验设计角度的检查。

4.1 常见泄漏源(按隐蔽程度排序)

极其隐蔽(几乎发现不了):

  • 用全量数据算标准化统计量(mean/std)
  • 用全量数据构建词表(会把测试集的词泄露到词表)
  • 去重时用了测试集(可能删掉了测试集里的重复样本,间接泄漏)
  • 时序数据随机切分
  • 同一个人的多张图散落在 train/val

比较隐蔽:

  • 数据增强在划分之前做(同一张图的增强版本跨集)
  • 早停用了测试集(应该用验证集)
  • 反复用测试集选模型

容易发现:

  • 训练集和测试集有完全相同的样本
  • 标签泄漏(特征里直接包含答案)

4.2 一个实用的检测方法

验证集性能异常好时,先假设有泄漏:

python
# 检测 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("⚠️ 打乱标签后仍有高准确率 → 泄漏!")
1
2
3
4
5
6
7
8
9
10

第二个检测极其有力(来自 Kaggle 的泄漏检测技巧):如果把标签完全打乱,模型还能得到高准确率,那就说明特征里直接包含答案。


五、实验的记录与复现

5.1 必须记录的东西

每个实验至少记录:

python
{
    "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,
}
1
2
3
4
5
6
7
8
9
10

最容易被忽略的是前两项(git commit 和数据版本):

  • 代码改了但不知道是哪个版本跑的结果 → 无法复现
  • 数据预处理改了但没记 → 对比失去意义

一个真实的惨痛案例:花两周调试一个改进,最后发现是「上一轮实验遗留的 model.eval() 没删」,之前所有结果都作废。这就是为什么要用配置管理。

5.2 推荐的工具

工具用途
git代码版本(必须)
W&B / MLflow / TensorBoard实验追踪
配置文件超参和实验逻辑
固定 seed可复现
Docker环境一致

最低配置(个人项目):

python
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 summary
1
2
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:统计显著性的正确做法

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

实测输出:

差值: 1.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才显著)
1
2
3
4

这个实验的核心发现:

  • 独立样本:SE = 0.80%,改进只有 1% → 只有 1.25 个 SE,完全不显著
  • 95% 置信区间 [0.8344, 0.8656] 里包含了 0.84(baseline 的性能)
  • 配对检验:χ2=0.450\chi^2 = 0.450χ2=0.450,远小于显著性阈值 3.84 → 更不显著

为什么配对检验更不显著? 因为配对检验消掉了「样本难度差异」这个大噪声源。如果 A 在困难样本上也赢、B 在简单样本上赢,配对检验能看清;独立检验会把这些「赢的地方」和「输的地方」混在一起算,反而高估显著性。

教训:在这个样本量下,1% 的改进根本无法可靠检测。 要可靠检测 1% 的差异,需要大约 4 倍的测试样本。

实验 2:消融实验的完整设计

python
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 不重要' 的错误结论")
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21

这个演示说明了消融实验最容易犯的错误:不做组合消融,会低估某些组件的重要性。

实验 3:数据泄漏检测

python
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 会让测试集信息泄漏到训练")
1
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 篇的过拟合完全同构,只是过拟合的对象从「训练数据」变成了「验证集」。

其他可能:

  1. 测试集分布和验证集不同 —— 检查数据划分是否合理
  2. 测试集太小 —— SE 太大,1% 差异看不出来
  3. 改进是数据特定的 —— 需要在更多数据集上验证
  4. 真的没有改进 —— 验证集的 1% 是噪声

诊断方法:

1. 用多个 seed 跑,看验证集提升是否稳定
2. 如果 std > 1%,那验证集的提升就是噪声
3. 检查是否试了太多组超参(>10组就要警惕)
1
2
3

根本的解法:用交叉验证(K 折),这样每个样本都当过验证集,能得到更可靠的估计。

一个实用建议:把测试集分成两半,一半用来做最终选择,一半做最终报告。这样能检测出「我是不是在验证集上过拟合了」。

Q3:一个论文报告「我们的方法在 ImageNet 上达到 85.0%,超过 ResNet-50 的 84.5%」。你怎么判断这个改进是否可信?

答案

按第 3.1 节的五步逐一检查:

① 差异是否显著? ImageNet val 有 50000 张图。

SE=0.85×0.1550000≈0.16%\text{SE} = \sqrt{\frac{0.85 \times 0.15}{50000}} \approx 0.16\% 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 文件,无网也能看

面试怎么用这两份

不要按顺序背。按题目反查:

  1. 打开面试高频问题,找你不确定的那几题
  2. 每道题后面都标了对应的正文篇目,点进去看推导
  3. 合上文档,用自己的话把这道题讲一遍
  4. 讲不出来的地方,回去看公式速查表

能讲清楚「这个方法什么时候会失效」,比背下公式更管用。

离线单页版

深度学习进阶手册.html 是全部内容的单文件打包版,公式和代码高亮都已渲染。 下载后双击就能看,不需要网络,也不需要跑这个站点。

公式速查表

所有关键公式一页汇总。适合复习和面试前速记。


优化

梯度下降

θt+1=θt−η∇θL\theta_{t+1} = \theta_t - \eta \nabla_\theta L θt+1​=θt​−η∇θ​L

SGD with Momentum

vt=βvt−1+gt,θt=θt−1−ηvtv_t = \beta v_{t-1} + g_t, \qquad \theta_t = \theta_{t-1} - \eta v_t vt​=βvt−1​+gt​,θt​=θt−1​−ηvt​

展开:vt=gt+βgt−1+β2gt−2+⋯v_t = g_t + \beta g_{t-1} + \beta^2 g_{t-2} + \cdotsvt​=gt​+βgt−1​+β2gt−2​+⋯(指数加权平均)

RMSProp

st=βst−1+(1−β)gt2,θt=θt−1−ηst+ϵgts_t = \beta s_{t-1} + (1-\beta)g_t^2, \qquad \theta_t = \theta_{t-1} - \frac{\eta}{\sqrt{s_t}+\epsilon}g_t st​=βst−1​+(1−β)gt2​,θt​=θt−1​−st​​+ϵη​gt​

Adam

mt=β1mt−1+(1−β1)gtvt=β2vt−1+(1−β2)gt2m^t=mt/(1−β1t)v^t=vt/(1−β2t)θt=θt−1−ηm^tv^t+ϵ\begin{aligned} m_t &= \beta_1 m_{t-1} + (1-\beta_1)g_t \\ v_t &= \beta_2 v_{t-1} + (1-\beta_2)g_t^2 \\ \hat{m}_t &= m_t / (1-\beta_1^t) \\ \hat{v}_t &= v_t / (1-\beta_2^t) \\ \theta_t &= \theta_{t-1} - \eta\frac{\hat{m}_t}{\sqrt{\hat{v}_t}+\epsilon} \end{aligned} mt​vt​m^t​v^t​θt​​=β1​mt−1​+(1−β1​)gt​=β2​vt−1​+(1−β2​)gt2​=mt​/(1−β1t​)=vt​/(1−β2t​)=θt−1​−ηv^t​​+ϵm^t​​​

AdamW(解耦 weight decay)

θt=θt−1−ηm^tv^t+ϵ−ηλθt−1\theta_t = \theta_{t-1} - \eta\frac{\hat{m}_t}{\sqrt{\hat{v}_t}+\epsilon} - \eta\lambda\theta_{t-1} θt​=θt−1​−ηv^t​​+ϵm^t​​−ηλθt−1​

LLM 标准超参:lr=1e-4∼3e-4\text{lr}=1\text{e-}4 \sim 3\text{e-}4lr=1e-4∼3e-4,β2=0.95\beta_2=0.95β2​=0.95,wd=0.1\text{wd}=0.1wd=0.1,clip=1.0,warmup 2000 步 + cosine decay


正则化

L2 / weight decay

θ←(1−ηλ)θ−η∂L∂θ\theta \leftarrow (1-\eta\lambda)\theta - \eta\frac{\partial L}{\partial\theta} θ←(1−ηλ)θ−η∂θ∂L​

bias 和 1 维参数(BN/LN 的 γ,β\gamma,\betaγ,β)不做衰减。

L1

L~=L+λ∑j∣θj∣\tilde{L} = L + \lambda\sum_j|\theta_j| L~=L+λj∑​∣θj​∣

产生稀疏性(不可导点在 0)。

Dropout

训练时以概率 ppp 置零,并除以 1−p1-p1−p;推理时不丢弃。


反向传播

三条基本规则

运算梯度
z=a+bz = a + bz=a+b∂L/∂a=∂L/∂z\partial L/\partial a = \partial L/\partial z∂L/∂a=∂L/∂z
z=abz = abz=ab∂L/∂a=(∂L/∂z)b\partial L/\partial a = (\partial L/\partial z)b∂L/∂a=(∂L/∂z)b
z=a⊤Wz = a^\top Wz=a⊤W∂L/∂a=W(∂L/∂z)⊤\partial L/\partial a = W(\partial L/\partial z)^\top∂L/∂a=W(∂L/∂z)⊤,∂L/∂W=a(∂L/∂z)⊤\partial L/\partial W = a(\partial L/\partial z)^\top∂L/∂W=a(∂L/∂z)⊤

梯度消失/爆炸的量级

∂L∂x1∝∏l=1L1n=n−L/2\frac{\partial L}{\partial x_1} \propto \prod_{l=1}^{L} \frac{1}{\sqrt{n}} = n^{-L/2} ∂x1​∂L​∝l=1∏L​n​1​=n−L/2

sigmoid:σ′(z)≤0.25\sigma'(z) \le 0.25σ′(z)≤0.25,每层乘 0.25


初始化

方法σw2\sigma_w^2σw2​适用
Xavier (Glorot)2nin+nout\frac{2}{n_{in}+n_{out}}nin​+nout​2​tanh / sigmoid
He (Kaiming)2nin\frac{2}{n_{in}}nin​2​ReLU / GELU
LLaMAN(0,0.022)N(0, 0.02^2)N(0,0.022)配合 RMSNorm

归一化

BatchNorm

x^=x−μBσB2+ϵ,y=γx^+β\hat{x} = \frac{x - \mu_B}{\sqrt{\sigma_B^2 + \epsilon}}, \qquad y = \gamma\hat{x} + \beta x^=σB2​+ϵ​x−μB​​,y=γx^+β

统计维度:batch + 空间 → 训练/推理行为不同

LayerNorm

x^=x−μσ2+ϵ\hat{x} = \frac{x - \mu}{\sqrt{\sigma^2+\epsilon}} x^=σ2+ϵ​x−μ​

统计维度:特征 → 与 batch 无关

RMSNorm

RMSNorm(x)=x1d∑jxj2+ϵ⊙γ\text{RMSNorm}(x) = \frac{x}{\sqrt{\frac{1}{d}\sum_j x_j^2 + \epsilon}} \odot \gamma RMSNorm(x)=d1​∑j​xj2​+ϵ​x​⊙γ

省掉减均值,只有 γ\gammaγ。


CNN

卷积参数量

params=k×k×Cin×Cout+Cout\text{params} = k \times k \times C_{in} \times C_{out} + C_{out} params=k×k×Cin​×Cout​+Cout​

感受野

rl=rl−1+(kl−1)∏i=1l−1si,jl=jl−1+(kl−1)∏i=1l−1sir_l = r_{l-1} + (k_l - 1)\prod_{i=1}^{l-1}s_i, \qquad j_l = j_{l-1} + (k_l - 1)\prod_{i=1}^{l-1}s_i rl​=rl−1​+(kl​−1)i=1∏l−1​si​,jl​=jl−1​+(kl​−1)i=1∏l−1​si​

stride=1 简化:rl=1+∑i=1l(ki−1)r_l = 1 + \sum_{i=1}^{l}(k_i - 1)rl​=1+∑i=1l​(ki​−1)

空洞卷积

有效感受野=k+(k−1)(d−1),需padding=d\text{有效感受野} = k + (k-1)(d-1), \quad \text{需padding} = d 有效感受野=k+(k−1)(d−1),需padding=d


ResNet

残差连接

y=F(x)+x,∂y∂x=∂F∂x+Iy = F(x) + x, \qquad \frac{\partial y}{\partial x} = \frac{\partial F}{\partial x} + I y=F(x)+x,∂x∂y​=∂x∂F​+I

梯度连乘展开:

∏l=1L(I+Jl)=I+∑Jl+∑l<kJlJk+⋯\prod_{l=1}^{L}(I + J_l) = I + \sum J_l + \sum_{l<k} J_lJ_k + \cdots l=1∏L​(I+Jl​)=I+∑Jl​+l<k∑​Jl​Jk​+⋯

展开后第一项是 III,所以梯度不衰减。


Attention

Scaled Dot-Product Attention

Attention(Q,K,V)=softmax(QK⊤dk+M)V\text{Attention}(Q,K,V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} + M\right)V Attention(Q,K,V)=softmax(dk​​QK⊤​+M)V

为什么除 dk\sqrt{d_k}dk​​:点积方差 Var[q⋅k]=dk\text{Var}[q\cdot k] = d_kVar[q⋅k]=dk​,std =dk=\sqrt{d_k}=dk​​。不缩放会让 softmax 饱和(实测 dk=64d_k=64dk​=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\text{MultiHead}(Q,K,V) = \text{Concat}(\text{head}_1,\ldots,\text{head}_h)W^O MultiHead(Q,K,V)=Concat(head1​,…,headh​)WO

参数量和单头相同:4dmodel24d_{model}^24dmodel2​


RoPE

θi=base−2i/d,Rmq=(cos⁡mθ−sin⁡mθsin⁡mθcos⁡mθ)q\theta_i = \text{base}^{-2i/d}, \qquad R_m q = \begin{pmatrix}\cos m\theta & -\sin m\theta \\ \sin m\theta & \cos m\theta\end{pmatrix}q θi​=base−2i/d,Rm​q=(cosmθsinmθ​−sinmθcosmθ​)q

核心性质:

⟨Rmq,Rnk⟩=⟨q,Rn−mk⟩=f(q,k,n−m)\langle R_m q, R_n k\rangle = \langle q, R_{n-m}k\rangle = f(q,k,n-m) ⟨Rm​q,Rn​k⟩=⟨q,Rn−m​k⟩=f(q,k,n−m)

内积只依赖相对位置。 base=10000(几何频率)。


现代 LLM 组件

SwiGLU

SwiGLU(x)=SiLU(xW1)⊗(xW3),SiLU(x)=xσ(x)\text{SwiGLU}(x) = \text{SiLU}(xW_1) \otimes (xW_3), \qquad \text{SiLU}(x) = x\sigma(x) SwiGLU(x)=SiLU(xW1​)⊗(xW3​),SiLU(x)=xσ(x)

三个矩阵,中间维度取 83d\frac{8}{3}d38​d 以保持参数量。

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

KV Cache

KV cache=2×nkv×L×dhead×nlayers×bytes\text{KV cache} = 2 \times n_{kv} \times L \times d_{head} \times n_{layers} \times \text{bytes} 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)=(NcN)αN,αN≈0.076\text{Kaplan:}\quad L(N) = \left(\frac{N_c}{N}\right)^{\alpha_N}, \quad \alpha_N \approx 0.076 Kaplan:L(N)=(NNc​​)αN​,αN​≈0.076

Chinchilla:Nopt∝C0.55,Dopt∝C0.45\text{Chinchilla:}\quad N_{opt} \propto C^{0.55}, \quad D_{opt} \propto C^{0.45} Chinchilla:Nopt​∝C0.55,Dopt​∝C0.45

Token/参数比最优 ≈ 20(但实践上为推理效率会超过这个值)。


损失函数

损失来源假设公式任务
MSE高斯1N∑(y−y^)2\frac{1}{N}\sum(y-\hat{y})^2N1​∑(y−y^​)2回归
MAE拉普拉斯$\frac{1}{N}\sumy-\hat
CrossEntropy类别−log⁡p^y-\log \hat{p}_y−logp^​y​多分类
BCE伯努利交叉熵(二分类)二分类
NLL—−∑log⁡p-\sum\log p−∑logp配合 LogSoftmax
InfoNCE对比−log⁡esim(i,i)∑jesim(i,j)-\log\frac{e^{sim(i,i)}}{\sum_j e^{sim(i,j)}}−log∑j​esim(i,j)esim(i,i)​对比学习

关键:CrossEntropyLoss / BCEWithLogitsLoss 接受 logits,不是概率。


统计

偏差-方差分解

E[(y−y^)2]=Bias2+Var+σ2\mathbb{E}[(y - \hat{y})^2] = \text{Bias}^2 + \text{Var} + \sigma^2 E[(y−y^​)2]=Bias2+Var+σ2

标准误与置信区间

SE=p(1−p)n,95%CI=p±1.96 SE\text{SE} = \sqrt{\frac{p(1-p)}{n}}, \qquad 95\%\text{CI} = p \pm 1.96\,\text{SE} SE=np(1−p)​​,95%CI=p±1.96SE

泛化差距

gap=train_loss−val_loss(或 train_acc−val_acc)\text{gap} = \text{train\_loss} - \text{val\_loss} \quad (\text{或 train\_acc} - \text{val\_acc}) gap=train_loss−val_loss(或 train_acc−val_acc)


数值稳定

Loss Scaling

实际梯度=S⋅真实梯度\text{实际梯度} = S \cdot \text{真实梯度} 实际梯度=S⋅真实梯度

bf16 不需要(指数位与 fp32 相同,动态范围 10±3810^{\pm38}10±38)。

梯度裁剪

g^=g⋅max_norm∥g∥当 ∥g∥>max_norm\hat{g} = g\cdot\frac{\text{max\_norm}}{\|g\|} \quad \text{当 } \|g\| > \text{max\_norm} g^​=g⋅∥g∥max_norm​当 ∥g∥>max_norm

等比缩放(不是逐元素裁剪,后者改变梯度方向)。


数值速查

概念数值
LLaMA-7B 参数量6.7B(transformer 部分 5.37B)
LLaMA-7B 隐藏维度4096
LLaMA-7B 头数 / dkd_kdk​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\lambda\|\theta\|^2λ∥θ∥2,梯度 λθ\lambda\thetaλθ → 参数按比例收缩,不会恰好为 0
  • L1 惩罚 λ∥θ∥1\lambda\|\theta\|_1λ∥θ∥1​,不可导点在 0 → 参数被压到恰好为 0
  • 深层原因:L1 对应拉普拉斯先验(密度在 0 处有尖峰),L2 对应高斯先验
  • 联系:w←(1−ηλ)w−η∇Lw \leftarrow (1-\eta\lambda)w - \eta\nabla Lw←(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篇

BatchNormLayerNorm
统计维度batch + 空间特征
依赖 batch是否
训练/推理行为不同相同

LLM 用 LN 的三个原因:batch 小(统计量不可靠)、变长序列(padding 污染)、推理需要确定性。


二、优化(7-13)

7. SGD 和 Adam 的区别?AdamW 又改进了什么? → Part1 第4篇

  • SGD:只有 lr,无自适应
  • Adam:一阶矩(动量)+ 二阶矩(自适应步长)
  • AdamW:解耦 weight decay。Adam 里衰减会进入 mt,vtm_t, v_tmt​,vt​,被 1vt\frac{1}{\sqrt{v_t}}vt​​1​ 缩放,行为不可预测;AdamW 直接 θ←θ−ηλθ\theta \leftarrow \theta - \eta\lambda\thetaθ←θ−ηλθ
  • LLM 为什么全用 AdamW:大数据下过拟合风险低,AdamW 更稳 + 可解耦其他衰减策略

8. 为什么 Adam 需要偏差修正? → Part1 第4篇

  • m0=0m_0=0m0​=0 → 第一步 m1=(1−β)g1=0.1g1m_1 = (1-\beta)g_1 = 0.1g_1m1​=(1−β)g1​=0.1g1​,严重低估
  • 修正:除以 1−βt1-\beta^t1−βt(几何级数的归一化因子)
  • 效果:第一步的有效步长精确等于名义 lr

9. 为什么 LLM 用 β2=0.95\beta_2 = 0.95β2​=0.95 而不是 0.999? → Part1 第4篇

  • 11−β2\frac{1}{1-\beta_2}1−β2​1​ = 二阶矩的有效窗口长度
  • 0.999 → 1000 步窗口,训练中梯度尺度快速变化时会严重滞后
  • 0.95 → 20 步窗口,能跟上当前状态
  • 这是 HuggingFace / LLaMA 的默认配置,已成为事实标准

10. 为什么需要 warmup? → Part1 第4篇

  • 初期梯度方向噪声大,直接用满 lr 会破坏已学到的结构
  • Adam 早期 mt,vtm_t, v_tmt​,vt​ 估计不准
  • 配合 weight decay:实际衰减率 = ηλ\eta\lambdaηλ,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±510^{\pm5}10±5,深层梯度 10−810^{-8}10−8 会下溢成 0
  • loss scaling:梯度放大 S=65536S=65536S=65536 倍,scaler.step() 内部除回去,动态调整 S(遇 inf/nan 就减小)
  • bf16 指数位与 fp32 相同(10±3810^{\pm38}10±38),不下溢,不需要 scaling
  • 注意:用梯度裁剪时必须先 scaler.unscale_(optimizer)

13. 学习率怎么调?太大太小分别什么现象? → Part1 第4篇、Part1 第6章

lr 过大lr 过小
loss 震荡 / 变 nanloss 几乎不降
训练像在发散收敛极慢
  • 最优 lr 随 batch size 缩放(线性缩放法则)
  • LLM 微调:1e-4 ~ 3e-4(AdamW);预训练:1e-4 左右

三、深度学习原理(14-20)

14. 反向传播的链式法则怎么用? → Part2 第5篇

三条规则:加法→等值分发;乘法→交叉相乘;矩阵乘→外积。

∂L∂W=a⊤⋅∂L∂z\frac{\partial L}{\partial W} = a^\top \cdot \frac{\partial L}{\partial z} ∂W∂L​=a⊤⋅∂z∂L​

踩坑点:忘了 batch 平均的系数、矩阵乘法顺序写反(PyTorch 会静默广播出错结果)。


15. 梯度消失的本质是什么?残差连接为什么有效? → Part2 第8篇

  • 本质:∂L∂x1∝∏l1n=n−L/2\frac{\partial L}{\partial x_1} \propto \prod_l \frac{1}{\sqrt{n}} = n^{-L/2}∂x1​∂L​∝∏l​n​1​=n−L/2,深层指数衰减
  • 残差的保证:∏l(I+Jl)\prod_{l}(I + J_l)∏l​(I+Jl​) 展开后第一项是 III,即使所有 Jl=0J_l = 0Jl​=0,梯度也原样传下去
  • 实测:60 层无残差梯度 = 0,有残差 = 1.01e+05

16. 为什么不能用 0 初始化权重? → Part2 第9篇

  • 全 0 → 输出 0 → 梯度 0 → 完全不训练
  • 对称性问题:所有神经元梯度相同 → 永远同步更新 → 等效于一个神经元
  • 例外:bias 可以初始化为 0;残差块最后一层的 BN 可以设 γ=0\gamma=0γ=0(让块初始为恒等)

17. Xavier 和 He 的区别?为什么 ReLU 网络要用 He? → Part2 第9篇

  • Xavier:σ2=2nin+nout\sigma^2 = \frac{2}{n_{in}+n_{out}}σ2=nin​+nout​2​,兼顾前向和反向方差保持,适合对称激活
  • He:σ2=2nin\sigma^2 = \frac{2}{n_{in}}σ2=nin​2​,补偿 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篇

  • 少了:求均值、减均值、β\betaβ 参数
  • 核心:归一化的主要作用是「控制尺度」,不是「控制中心」
  • 实测:RMSNorm 保留输入偏移(+0.426),但标准差一样被控制(0.905)
  • 收益在算子融合,不在参数(省 ddd 个参数微不足道)

20. 什么是「退化问题」?和过拟合有什么区别? → Part2 第8篇

  • 退化:网络加深,训练误差也上升 → 是优化问题,不是过拟合
  • 过拟合:训练误差低,验证误差高
  • 原因:优化器难以在深层网络中「找到」恒等映射的解

四、Transformer(21-27)

21. Attention 的公式?为什么除以 dk\sqrt{d_k}dk​​? → Part3 第10篇

Attention=softmax(QK⊤dk+M)V\text{Attention} = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}} + M\right)V Attention=softmax(dk​​QK⊤​+M)V

为什么:Var[q⋅k]=dk\text{Var}[q\cdot k] = d_kVar[q⋅k]=dk​ → std = dk\sqrt{d_k}dk​​。不缩放 → logits std 太大 → softmax 饱和。

实测:dk=64d_k=64dk​=64 时不缩放 max_p = 0.9990(完全 one-hot,梯度消失),缩放后 0.0350。


22. 多头注意力的参数量比单头多吗? → Part3 第11篇

一样多,都是 4dmodel24d_{model}^24dmodel2​。多头是把一个大注意力拆成 hhh 个小的,不是增加参数。

「多」的体现:hhh 个独立的注意力图,能学到不同的关系模式。


23. RoPE 为什么比绝对位置编码好? → Part3 第11篇

核心性质:⟨Rmq,Rnk⟩=f(q,k,n−m)\langle R_m q, R_n k\rangle = f(q,k,n-m)⟨Rm​q,Rn​k⟩=f(q,k,n−m),内积只依赖相对位置。

四个优势:内置相对位置关系 / 不占表示维度 / 外推更好 / 是正交变换(不破坏语义结构)。

实测:同一 (q,k)(q,k)(q,k) 放不同位置,内积波动仅 1e-6。


24. 什么是 KV Cache?为什么能加速?代价是什么? → Part3 第13篇

  • 原理:causal mask 保证已生成 token 的 kj,vjk_j, v_jkj​,vj​ 永不改变 → 存下来复用
  • 收益:计算量 O(n3)→O(n2)O(n^3) \to O(n^2)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 的比例)
MHAhhh100%
GQAgggg/hg/hg/h
MQA13%

论文结论:h/g=4∼8h/g = 4\sim8h/g=4∼8 时质量接近 MHA。实践:LLaMA3-70B 用 g=8g=8g=8,7B 用 g=8g=8g=8(h/g=4h/g=4h/g=4)。


26. Prefill 和 Decode 的瓶颈有什么不同?优化手段? → Part3 第13篇

PrefillDecode
瓶颈算力显存带宽
矩阵形状[2048,d]×[d,d][2048, d]\times[d,d][2048,d]×[d,d][1,d]×[d,d][1, d]\times[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 → 残差
1
2

必备组件:RMSNorm、MultiHeadAttention(含 RoPE)、SwiGLU、causal mask、final norm。

参数量配置:4 个 d2d^2d2(QKVO)+ 3×d×hffn3 \times d \times h_{ffn}3×d×hffn​(SwiGLU,hffn≈83dh_{ffn} \approx \frac{8}{3}dhffn​≈38​d)。


五、研究方法(28-30)

28. 论文报告的提升 1%,可信吗?怎么验证? → Part4 第16篇

五步验证:

  1. 算置信区间(n=1000n=1000n=1000 时 1% 差异可能只有 1.4 个 SE)
  2. 3+ 随机种子报 ±std
  3. 配对检验(McNemar),不是独立 t 检验
  4. 给 baseline 也调参(最容易出问题的地方)
  5. 换数据集验证 + 报告代价

29. 消融实验怎么做才可信? → Part4 第16篇

四个原则:

  1. 一次只改一个东西
  2. 去掉后要重新调超参(最常被忽略)
  3. 多 seed
  4. 要做组合消融(检测交互作用)

组合消融的例子:单独去 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/inf8%清洗
梯度爆炸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 个错例(最有效的排查手段)


回到 目录