Stephen 技术博客

AI和数据平台的工程实践

← 返回文章列表

BatchNorm 反向传播完整推导

Makemore Part 4 专题讲义 · 分步七连击对应视频 Exercise 1 的 BN 段(约 50:00 – 1:26:31,Day 17)· 融合公式对应 Exercise 3(1:36:37 – 1:50:02,Day 18)
配套练习:Makemore_Part4_实战_BN反向传播_练习.ipynb(推荐,一步一个 cell)· 另有 .py 脚本版
🐉 专治大魔王

一、为什么这 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:

x μ (mean) σ² (var) x̂ (bnraw) y → loss μ 也进 σ² ① 直连 ② 经 μ ③ 经 σ²

反向传播时,三条支路的梯度必须全部算出来再相加(Micrograd 的老规矩:分叉汇合处梯度累加)。难点全在这:漏一条支路、或忘了 broadcast 的 .sum(0),数值就对不上。

破解策略:不要直接推三条支路的总和!把 forward 拆成 7 个"傻瓜算子",逐个反向,每步 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] 的形状。)

最容易错的两个地方:STEP 5 和 STEP 7 的支路汇合 —— bndiffhprebn 各有两条支路,必须相加而不是只算一条或互相覆盖。漏掉任何一条,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 求导:

∂L/∂σ² = Σᵢ γgᵢ · (xᵢ-μ) · (-½)·v^(-3/2) = -½·v^(-1)·Σᵢ γgᵢ·x̂ᵢ

(第二个等号用了 (xᵢ-μ)·v^(-3/2) = x̂ᵢ·v^(-1),这种"用 forward 缓存变量改写"是化简的核心技巧。)

第 2 步 —— 对 μ 的梯度(支路②的源头)。μ 通过两条路影响 loss:直接出现在 x̂ 里,以及出现在 σ² 里。

∂L/∂μ = Σᵢ γgᵢ·(-v^(-1/2)) + ∂L/∂σ² · ∂σ²/∂μ

关键观察:∂σ²/∂μ = -2/(n-1)·Σᵢ(xᵢ-μ) = 0,因为偏差和恒等于零(Σ(xᵢ-μ) = Σxᵢ - n·μ = 0)。所以第二项直接消失:

∂L/∂μ = -v^(-1/2) · Σᵢ γgᵢ
这就是 Karpathy 纸上推导中"莫名其妙消掉一大坨"的那一步:σ² 对 μ 的敏感度为 0,不是巧合,是"均值附近的偏差和为零"这个恒等式。很多人卡死在这里以为自己算错了。

第 3 步 —— 汇总到 xⱼ。xⱼ 的三条支路:直连(系数 v^(-1/2))、经 σ²(∂σ²/∂xⱼ = 2(xⱼ-μ)/(n-1))、经 μ(∂μ/∂xⱼ = 1/n):

∂L/∂xⱼ = γgⱼ·v^(-1/2) + ∂L/∂σ²·2(xⱼ-μ)/(n-1) + ∂L/∂μ·(1/n)

代入第 1、2 步的结果,并把 (xⱼ-μ) 改写成 x̂ⱼ·v^(1/2):

= γv^(-1/2)·[ gⱼ − (1/n)·Σᵢgᵢ − (1/(n-1))·x̂ⱼ·Σᵢgᵢx̂ᵢ ]

提出 1/n,就是 Karpathy 的最终形式:

∂L/∂xⱼ = (γ·v^(-1/2)/n) · [ n·gⱼ − Σᵢgᵢ − (n/(n-1))·x̂ⱼ·Σᵢgᵢx̂ᵢ ]

翻译成 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 报错时按顺序查)

  1. 形状先行dbngain.shape == bngain.shape == [1, 64]?忘了 keepdim=True 会变成 [64],后面全错。
  2. 广播必 sum:forward 里被广播复用的变量(bngain、bnbias、bnvar_inv、bnmeani),backward 一律 .sum(0, keepdim=True)
  3. 支路必相加:bndiff(STEP 2 + STEP 5)、hprebn(STEP 6 + STEP 7)各有两条支路,汇合时 支路1 + 支路2,一条都不能漏。
  4. notebook 特有坑:如果你自己写了 +=,检查是不是重复运行过那个 cell(重复累加)。稳妥做法:从头 Restart & Run All 一遍再看 cmp 结果。
  5. n-1 还是 n:STEP 4 的系数必须和 forward 的 bnvar 一致。
  6. ε 别丢:STEP 3 求导时里面是 (bnvar + 1e-5),不是裸 bnvar。
  7. 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 都有同款"自身项 − 聚合修正项"结构。这次推透了,以后见到都是老朋友。