BatchNorm 反向传播完整推导
一、为什么这 20 分钟这么难:三条支路
先看清楚敌人。BN 的 forward(记 x = hprebn,形状 [n, 64],n = 32):
μ = (1/n) · Σᵢ xᵢ # bnmeani [1, 64]
σ² = (1/(n-1)) · Σᵢ (xᵢ - μ)² # bnvar [1, 64](无偏!)
x̂ = (x - μ) / √(σ² + ε) # bnraw [n, 64]
y = γ·x̂ + β # hpreact [n, 64]
普通算子(tanh、矩阵乘)的梯度只沿一条边回流。而 BN 里,x 通过三条路径影响 loss:
反向传播时,三条支路的梯度必须全部算出来再相加(Micrograd 的老规矩:分叉汇合处梯度累加)。难点全在这:漏一条支路、或忘了 broadcast 的 .sum(0),数值就对不上。
cmp() 验证。每一步都只用你在 Micrograd / Exercise 1 里已经会的规则。练习 notebook 就是按这个思路组织的:一步一个 cell,填完 TODO 按 Shift+Enter 立刻验证,不用整体重跑。
二、七步分解:每步只是一个小算子
| 步 | Forward(倒序反推) | Backward | 用到的规则 |
|---|---|---|---|
| 1 | hpreact = bngain*bnraw + bnbias |
dbngain = (dhpreact*bnraw).sum(0,kd)dbnbias = dhpreact.sum(0,kd)dbnraw = dhpreact*bngain |
乘法 + 加法;γ、β 是 [1,64] 被广播 → sum(0) |
| 2 | bnraw = bndiff * bnvar_inv |
dbnvar_inv = (dbnraw*bndiff).sum(0,kd)dbndiff_b1 = dbnraw*bnvar_inv ←仅支路①,先存着 |
乘法;bnvar_inv 被广播 → sum(0) |
| 3 | bnvar_inv = (bnvar+ε)**-0.5 |
dbnvar = -0.5*(bnvar+ε)**-1.5 * dbnvar_inv |
幂函数 d(uᵃ)=a·uᵃ⁻¹ |
| 4 | bnvar = (1/(n-1))*bndiff2.sum(0,kd) |
dbndiff2 = (1/(n-1))*ones_like(bndiff2)*dbnvar |
sum 的反向 = 复制扩展 |
| 5 | bndiff2 = bndiff**2 |
dbndiff = dbndiff_b1 + 2*bndiff*dbndiff2 ←支路③并入 |
平方求导;支路相加 |
| 6 | bndiff = hprebn - bnmeani |
dbnmeani = (-dbndiff).sum(0,kd)dhprebn_b1 = dbndiff ←仅支路①,先存着 |
减法;bnmeani 被广播 → sum(0),带负号 |
| 7 | bnmeani = (1/n)*hprebn.sum(0,kd) |
dhprebn = dhprebn_b1 + (1/n)*ones_like(hprebn)*dbnmeani ←支路②并入 |
sum 的反向 = 复制扩展;支路相加 |
(kd = keepdim=True,保住 [1, 64] 的形状。)
bndiff 和 hprebn 各有两条支路,必须相加而不是只算一条或互相覆盖。漏掉任何一条,cmp() 会给你 approximate: False。写法说明:notebook 版把支路 1 存成
dbndiff_b1/dhprebn_b1,汇合时写 支路1 + 支路2 —— 每个 cell 可以放心重跑;.py 脚本版(以及 Karpathy 视频里)用的是 += 累加写法,数学上完全等价。但在 notebook 里千万别用 +=:重复运行 cell 会重复累加,梯度越加越大,cmp 会莫名其妙报错。三、融合公式的完整数学推导
七步版每步都对之后,才有资格看这个。目标:把 7 步合并成一个式子。以下对某一列(某个 hidden 单元)推导,记:
gᵢ = dhpreact[i]—— 上游梯度(第 i 个样本)v = σ² + ε,所以bnvar_inv = v^(-1/2)x̂ᵢ = bnraw[i] = (xᵢ - μ)·v^(-1/2)
第 0 步:y = γ·x̂ + β,所以传到 x̂ 的梯度是 ∂L/∂x̂ᵢ = γ·gᵢ。下面全程带着 γ。
第 1 步 —— 对 σ² 的梯度(支路③的源头)。x̂ᵢ = (xᵢ-μ)·v^(-1/2),对 v 求导:
(第二个等号用了 (xᵢ-μ)·v^(-3/2) = x̂ᵢ·v^(-1),这种"用 forward 缓存变量改写"是化简的核心技巧。)
第 2 步 —— 对 μ 的梯度(支路②的源头)。μ 通过两条路影响 loss:直接出现在 x̂ 里,以及出现在 σ² 里。
关键观察:∂σ²/∂μ = -2/(n-1)·Σᵢ(xᵢ-μ) = 0,因为偏差和恒等于零(Σ(xᵢ-μ) = Σxᵢ - n·μ = 0)。所以第二项直接消失:
第 3 步 —— 汇总到 xⱼ。xⱼ 的三条支路:直连(系数 v^(-1/2))、经 σ²(∂σ²/∂xⱼ = 2(xⱼ-μ)/(n-1))、经 μ(∂μ/∂xⱼ = 1/n):
代入第 1、2 步的结果,并把 (xⱼ-μ) 改写成 x̂ⱼ·v^(1/2):
提出 1/n,就是 Karpathy 的最终形式:
翻译成 PyTorch(对所有 64 列并行,就是把 Σ 换成 .sum(0)):
dhprebn = (bngain * bnvar_inv / n) * (
n * dhpreact
- dhpreact.sum(0)
- (n/(n-1)) * bnraw * (dhpreact * bnraw).sum(0)
)
三项各自的物理意义
| 项 | 来自哪条支路 | 直觉 |
|---|---|---|
n·dhpreact | ① 直连 | 本样本自己的梯度,原样通过 |
- dhpreact.sum(0) | ② 经 μ | 你变大会拉高全 batch 的均值,"连累"其他样本 → 扣掉批平均梯度 |
- n/(n-1)·bnraw·(dhpreact·bnraw).sum(0) | ③ 经 σ² | 你偏离均值会撑大方差、压扁所有人的 x̂ → 按你的偏离程度 x̂ⱼ 扣除 |
一句话总结:BN 的梯度 = 你自己的梯度,减去"批均值修正",再减去"批方差修正"。这正是"每个样本的梯度依赖 batch 内所有其他样本"的数学表达 —— 也是 BN 和 LayerNorm 最本质的区别。
四、n/(n-1) 到底哪来的
唯一的来源:forward 里方差用的是无偏估计(Bessel 校正),分母是 n-1 而不是 n:
bnvar = (1/(n-1)) * bndiff2.sum(0, keepdim=True) # 对应 torch.var(unbiased=True) 默认行为
推导第 3 步里 ∂σ²/∂xⱼ = 2(xⱼ-μ)/(n-1) 的 n-1 就是它。合并化简、提出 1/n 之后,1/(n-1) 就变成了 n/(n-1)。如果 forward 用有偏方差(除 n),这个因子就是 1,公式更干净——很多教科书版本的 BN 反向公式没有这个因子,就是因为它们除的是 n。看别人的推导对不上时,先查这一点。
五、调试检查清单(cmp 报错时按顺序查)
- 形状先行:
dbngain.shape == bngain.shape == [1, 64]?忘了keepdim=True会变成 [64],后面全错。 - 广播必 sum:forward 里被广播复用的变量(bngain、bnbias、bnvar_inv、bnmeani),backward 一律
.sum(0, keepdim=True)。 - 支路必相加:bndiff(STEP 2 + STEP 5)、hprebn(STEP 6 + STEP 7)各有两条支路,汇合时
支路1 + 支路2,一条都不能漏。 - notebook 特有坑:如果你自己写了
+=,检查是不是重复运行过那个 cell(重复累加)。稳妥做法:从头 Restart & Run All 一遍再看 cmp 结果。 - n-1 还是 n:STEP 4 的系数必须和 forward 的 bnvar 一致。
- ε 别丢:STEP 3 求导时里面是
(bnvar + 1e-5),不是裸 bnvar。 - maxdiff 量级:1e-9 级 = 过关;1e-3 级 = 公式错了,不是浮点误差。
六、通关后:连回 1B 主线
- 为什么 LLM 都用 LayerNorm:LN 在特征维(而非 batch 维)归一化,样本之间零耦合 → 反向没有"连累项",公式简单得多,且与 batch size 无关、可流式推理。你刚推完 BN,现在能用 10 分钟推出 LN 的 backward 作为自测(提示:结构同第三节,把 Σᵢ 换到特征维,且没有跨样本耦合)。
- batch_size=1 时 BN 会怎样:n-1 = 0,方差除零直接爆炸 —— 这就是 BN 不能用于小 batch / 在线推理的根因之一。
- 三项结构会再次出现:Flash Attention 的 softmax backward、RMSNorm 的 backward 都有同款"自身项 − 聚合修正项"结构。这次推透了,以后见到都是老朋友。