Stephen 技术博客

AI和数据平台的工程实践

← 返回文章列表

Makemore Part 4 讲义:Backprop Ninja

对应 Karpathy "Becoming a Backprop Ninja"
视频链接:https://www.youtube.com/watch?v=q8SA3rM6ckI · 时长约 1h 56min
学习日:Day 16–19
⚡ 整个 30 天计划最硬的一节 ⚡

〇、本节的位置与意义

Karpathy 在视频开头亲口说:

"This is the hardest of the videos. I will not blame you if you skip it. But you should NOT skip it if you want to truly understand neural networks."

这是整个系列的分水岭:跳过去的人将永远停留在"PyTorch API 调用者"层级,看 Flash Attention 的论文像看天书,调梯度问题只能靠运气。

本节做什么:禁用 PyTorch 的 .backward(),手写每一个算子的反向传播,并用 cmp() 跟 autograd 对比,误差 < 1e-9 才算过关。

学完后,你拥有的四个"绝技":

绝技什么意思
梯度调试训练 NaN / 梯度爆炸时,能定位到具体哪一层、哪个操作出问题
论文阅读看 Flash Attention / LoRA / RLHF 里的 ∂L/∂x 公式不发懵
算子优化能写融合的 CUDA kernel,进入工业级训练的门槛
跟 PyTorch 平视.backward() 不再是黑魔法,你知道它每一步在干嘛

一、学习目标(Done when 清单)

学完后必须能:

  • □ 不看视频,手推 y = x @ W + b 的 backward(dx, dW, db)
  • □ 手推 tanh 的 backward,并解释为什么用 y 而不是 x 算导数
  • 在 30 秒内写出 softmax + cross_entropy 的简化 backward(3 行代码)
  • □ 解释为什么 cross_entropy 的简化公式 (probs - one_hot)/N 那么优雅
  • □ 手推 BatchNorm 的 backward(不要求一次写对,但能在 30 分钟内推导出来)
  • □ 解释 embedding lookup 的 backward 为什么需要 index_add_
  • □ 列出 reshape / view 操作的 backward 是什么
  • □ 解释 broadcasting 的 backward 为什么需要 .sum(0)
  • □ 用纯手动 backward 完整训练一个 MLP,loss 跟 autograd 版本一致
  • 核心检验:拿到一个新算子(如 GELU),10 分钟内推出它的 backward 公式

二、视频结构地图

视频约 1h 56min,分为 4 个练习:

练习时段主题难度
引言 + 起步代码0:00 – 13:01回顾 Part 3 的网络,介绍 cmp() 验证模式
Exercise 113:01 – 1:26:31手动 backprop 整个 forward 图(20+ 个变量,含分步 BN)——全片最长的一段★★★
插曲1:05:17Bessel 校正:方差为什么除 n−1(在 Ex1 中段出现)★★
Exercise 21:26:31 – 1:36:37cross_entropy 的优雅简化(一行搞定)★★★★
Exercise 31:36:37 – 1:50:02BatchNorm 的反向(视频最难的 20 分钟)★★★★★
Exercise 41:50:02 – 结尾把所有手动 backward 拼起来训练完整网络★★★

三、必备数学预备

看视频之前,下面这三件事必须在脑子里清晰:

3.1 链式法则的张量形式

标量版(Micrograd 里学过的):

dL/dx = dL/dy × dy/dx

张量版本:如果 y = f(x),其中 x、y 是张量,那么:

  • 给我 dout = dL/dy(上游传来的梯度,形状跟 y 相同)
  • 我返回 dx = dL/dx(要传给下游的梯度,形状跟 x 相同)
  • 关键规则:梯度的形状永远跟原变量相同

这是检查你公式对不对的"形状大法":算完后 dx.shape == x.shape,否则一定错了。

3.2 矩阵乘法的反向传播(必背)

这是 Part 4 最常用的一组公式:

forward:    y = x @ W

backward:   dx = dout @ W.T
            dW = x.T @ dout

记忆口诀

  • "反向乘以转置":dx 那一侧的 W 要转置
  • "形状会告诉你顺序":例如 W 是 [in, out],dW 必须是 [in, out],所以 x.T (in × N) @ dout (N × out) = dW (in × out)
验证小窍门:写完 backward 后,先用形状检查一遍,再用 cmp() 验证数值。形状不对的话,数值也一定不对。

3.3 广播 (broadcasting) 的反向

当 forward 里有广播,backward 需要"对广播的维度求和"。

forward:    y = x + b       # x: [N, D], b: [D] → 广播到 [N, D]

backward:   dx = dout            # 不需要做任何事
            db = dout.sum(0)     # 在 N 维度上求和

为什么?因为 b 被复用了 N 次,每次都参与了 forward,所以反向时每个样本的梯度都要加到 b 上。

规则:forward 时哪个维度被"复制扩展"了,backward 时对那个维度求和

3.4 reshape / view 的反向

forward:    y = x.view(-1, 30)    # x: [32, 3, 10] → y: [32, 30]

backward:   dx = dout.view(x.shape)    # 反向 reshape 回去

因为 view 不改变值,只改变形状,所以梯度也只需要 reshape 回去,不需要算任何数值。

四、四个练习详解

Exercise 1:完整 backward 网络

目标:把 Part 3 的网络 forward 拆成 20+ 个中间变量,每一个都手动算 backward,跟 PyTorch 的 .backward() 比对。

forward 大致长这样(伪代码):

# 11 个中间变量都要算 backward
emb     = C[Xb]                              # → demb
embcat  = emb.view(emb.shape[0], -1)         # → dembcat
hprebn  = embcat @ W1 + b1                   # → dhprebn
bnmeani = (1/n) * hprebn.sum(0, keepdim=True)
bndiff  = hprebn - bnmeani
bndiff2 = bndiff**2
bnvar   = (1/(n-1)) * bndiff2.sum(0, keepdim=True)
bnvar_inv = (bnvar + 1e-5)**-0.5
bnraw   = bndiff * bnvar_inv
hpreact = bngain * bnraw + bnbias
h       = torch.tanh(hpreact)
logits  = h @ W2 + b2
loss    = F.cross_entropy(logits, Yb)

Karpathy 的策略:从 loss 开始,倒着一个一个推。每推完一个就用 cmp() 验证。

难点

  • broadcasting 多次出现,要小心加 .sum(0)
  • view 操作的反向 = reshape 回去
  • 嵌入查表的反向 = 累加到对应行(用 index_add_
  • 一旦某一步错了,后面全部错 —— cmp() 帮你定位

预期时长:约 4 小时(仅视频段就有 73 分钟,含暂停推导)。Day 16–17 的主要内容。

Exercise 2:cross_entropy 的优雅简化

正常写 cross_entropy 的反向会非常繁琐:

logits → counts = exp(logits) → probs = counts / counts.sum() → loss = -log(probs[y]).mean()
# 反向需要分别推 log、divide、sum、exp 各一遍

但化简后只有 3 行

dlogits = F.softmax(logits, dim=1)    # 算 softmax 概率
dlogits[range(n), Yb] -= 1            # 正确类别的概率 -1
dlogits /= n                          # 平均

意义:这一段是整个视频最让人"啊哈"的时刻。它告诉你:

  • 数学上看似复杂的东西,化简后可能异常优雅
  • 这种简化在数值上也更稳定(避免了 log(0) 等)
  • 这就是为什么 PyTorch 的 F.cross_entropy 接受 logits 而非 probs
Karpathy 的原话:"This is the spirit of Backprop Ninja —— complex things become beautifully simple when you derive them carefully."

预期时长:30~40 分钟。Day 18 上半场。

Exercise 3:BatchNorm 的反向★★★★★ 大魔王

这是整个视频最难的 20 分钟。BatchNorm forward 看起来不复杂:

mu  = x.mean(0)
var = x.var(0)
xhat = (x - mu) / sqrt(var + eps)
y = gamma * xhat + beta

meanvar 是"聚合操作",反向传播时梯度要"分回去给每个样本",公式变得超级复杂:

dhprebn = (bngain * bnvar_inv / n) * (
    n * dhpreact
    - dhpreact.sum(0)
    - n/(n-1) * bnraw * (dhpreact * bnraw).sum(0)
)

5 个张量操作压缩在一行里。Karpathy 用 20 分钟在纸上推这一个公式。

⚠️ 心理准备:第一遍可能完全跟不上推导。这是正常的。建议:第一遍看完不强求懂,第二遍跟着暂停纸笔推导,第三遍能复述大意即可。

关键洞察:BatchNorm 的反向之所以复杂,是因为 forward 里每个样本都参与了 meanvar 的计算,所以每个样本的梯度都依赖于 batch 内所有其他样本。这是 BN 跟 LayerNorm 最大的区别之一。

预期时长:1~2 小时。Day 18 下半场(分步版已在 Exercise 1 里做过,这里攻融合公式)。

Exercise 4:拼起来训练

把前 3 个练习的 backward 函数全部串起来,完整训练一个 MLP,不用 loss.backward()

训练循环大致这样:

for step in range(max_steps):
    # forward
    ...

    # 手动 backward (替代 loss.backward())
    dlogits = ...
    dh, dW2, db2 = ...
    dhpreact = ...
    dbngain, dbnraw, dbnbias = ...
    dhprebn = ...
    dembcat, dW1, db1 = ...
    demb = ...
    dC = ...

    # update
    for p, grad in zip(parameters, grads):
        p.data += -lr * grad

    # 验证:手动 grad 应该跟 autograd 一致
    for p, grad in zip(parameters, grads):
        cmp(p, grad)

成就感时刻:当你看到训练完成后 val loss 跟 Part 3 一模一样,你就完成了"成为 backprop ninja"的全部仪式。

预期时长:1.5 小时。Day 19 的主要内容。

五、Backward 公式速查表

这张表在你写 Exercise 1 时不停翻:

Forward 操作Backward 公式说明
y = a + b da = dout
db = dout
有 broadcast 时记得 .sum(0)
y = a * b
(element-wise)
da = b * dout
db = a * dout
注意是 element-wise,不是矩阵乘
y = x @ W dx = dout @ W.T
dW = x.T @ dout
最重要的一对
y = tanh(x) dx = (1 - y**2) * dout 用 y 比用 x 算更省一次 tanh
y = exp(x) dx = y * dout 用 y = exp(x) 本身做导数
y = log(x) dx = (1/x) * dout 不能用 y,必须用 x
y = x.mean(0) dx = dout.repeat(N,1) / N 梯度均分给每个样本
y = x.sum(0) dx = dout.repeat(N,1) 每个样本拿一份
y = x.view(...) dx = dout.view(x.shape) 只 reshape,不算数
emb = C[X]
(embedding lookup)
dC.index_add_(0, X, demb) 梯度累加到被选中的行
softmax + cross_entropy
combined
dlogits = softmax(logits)
dlogits[range(n), y] -= 1
dlogits /= n
Exercise 2 的"优雅简化"
BatchNorm
combined
见 Exercise 3,5 个操作压缩在一行 Exercise 3 的"大魔王"

六、值得记住的几个"坑"

  1. 梯度形状 vs 原变量形状必须严格一致
    db.shape == b.shapedW.shape == W.shape。形状不对 = 数值一定不对。

  2. 有 broadcast 就要 .sum()
    比如 y = x + b,b 是 [D] 而 y 是 [N, D],反向时 db = dout.sum(0),绝不能只是 db = dout

  3. += 还是 =
    正在初始化的 grad 用 =,要"累加上一层传来的梯度"时用 +=。Micrograd 时代说过的规则。

  4. tanh 的导数用 y 而非 x
    公式 dx = (1 - y**2) * dout,其中 y = tanh(x)。用 y 省一次 tanh 计算(forward 时已经算过 y 了)。

  5. cross_entropy 的简化版要除以 N
    公式 dlogits /= n 这一行**不能忘**。因为原始 loss 是 .mean() 而非 .sum()

  6. embedding 的 backward 必须用 index_add_ 不能用赋值
    因为同一个字符可能在 batch 内出现多次,梯度要累加而非覆盖

  7. BatchNorm 的 n / (n-1) 因子很关键
    PyTorch 的 .var() 默认 unbiased=True,分母是 n-1 而非 n。反向公式要带这个因子。

  8. 验证用 cmp(),maxdiff < 1e-9 才算过关
    浮点误差容许,但应该在 1e-9 量级。如果 1e-3 那种"差不多",肯定是公式错了。

七、cmp() 验证模式

这是 Karpathy 在整节视频里反复用的"测试驱动开发"模式:

def cmp(s, dt, t):
    """
    s:  变量名(字符串,用于打印)
    dt: 你手算的梯度
    t:  PyTorch 算的版本(用 t.grad 拿)
    """
    ex = torch.all(dt == t.grad).item()
    app = torch.allclose(dt, t.grad)
    maxdiff = (dt - t.grad).abs().max().item()
    print(f'{s:15s} | exact: {str(ex):5s} | '
          f'approximate: {str(app):5s} | maxdiff: {maxdiff}')

使用方式

cmp('logits', dlogits, logits)
cmp('h', dh, h)
cmp('W2', dW2, W2)
# ...

理想输出

logits  | exact: False | approximate: True | maxdiff: 5.96e-09
h       | exact: False | approximate: True | maxdiff: 1.42e-08
W2      | exact: False | approximate: True | maxdiff: 3.73e-09
  • exact: True = 完全相等(很少能做到)
  • approximate: True + maxdiff < 1e-7 = 过关
  • approximate: False = 公式错了,回去检查

八、与 LLM / Transformer 的关系

这一节做的事情看起来很"原始",但其实是所有现代 LLM 训练栈的基础:

Part 4 学的对应到 1B+ 模型的什么
手动 backward 每个算子FlashAttention、xFormers 自己写的 attention backward
cross_entropy 的简化所有 LLM 训练 loss 函数的稳定实现
BatchNorm 反向的复杂度为什么 LLM 转用 LayerNorm(反向简单 1000 倍)
embedding lookup 反向token embedding 的梯度计算(GPT、LLaMA 都用)
cmp() 验证模式所有自定义 CUDA kernel 必须做的 gradient check
梯度形状必须一致分布式训练(FSDP、Megatron)的梯度通信基础

"手动 backward" 是工业级深度学习的"水电煤":DeepSpeed、Megatron、Triton 这些框架的核心就是高效的 backward 实现。学完 Part 4 你才有资格读它们的源码。

九、看视频前的预备问题

带着这些问题进入视频,效率比被动看高 2 倍:

Exercise 1 篇

  1. 线性层 y = x @ W + b 的 backward 三个公式分别是?
  2. 为什么 db = dout.sum(0) 不是 db = dout
  3. tanh 的 backward 用 y 还是 x 算?为什么用 y 更聪明?
  4. view 操作的 backward 是什么?需不需要算?

Exercise 2 篇

  1. cross_entropy 的简化公式 (probs - one_hot) / N 是怎么推出来的?
  2. 为什么这个简化在数值上比 "softmax → log → NLL" 更稳定?
  3. 这跟 F.cross_entropy(logits, y) 内部的实现一样吗?

Exercise 3 篇

  1. BatchNorm 反向为什么比一般算子复杂这么多?
  2. BN 公式里 n/(n-1) 这个因子从哪来的?
  3. 如果 batch_size = 1,BN backward 会出什么问题?

Exercise 4 篇

  1. 手动 backward 训练出的模型,val loss 应该跟 Part 3 完全一致还是仅近似?
  2. 为什么手动 backward 比 loss.backward() 慢得多?

十、看完后的自我检测

这 10 道题答不上来 3 道以上,回去重看:

#题目标准答案要点
1写出 y = x @ W + b 的三个 backward 公式dx = dout @ W.T, dW = x.T @ dout, db = dout.sum(0)
2tanh backward 的公式dx = (1 - y**2) * dout,y 是 tanh 输出
3cross_entropy 简化 backward 三行代码dlogits = softmax(logits); dlogits[range(n),y] -= 1; dlogits /= n
4broadcasting 加法的 backward 为什么要 sum因为 b 被复制 N 次,每次都参与 forward
5embedding 反向用 index_add_ 不用赋值的原因同一字符可能出现多次,梯度要累加
6view 的 backward 是什么dx = dout.view(x.shape),不算数
7BatchNorm 反向用 mean 的部分对应公式哪一项-dhpreact.sum(0) 这一项
8cmp() 验证的 maxdiff 多大算通过< 1e-7(理想 1e-9),表示数值上等价
9为什么手动 backward 是"训练 1B 模型的基础"分布式、混合精度、CUDA 算子都基于此
10核心检验:给你一个新算子 GELU,10 分钟内推出 backward能用链式法则 + 形状检查推出来即可

十一、Day 16–19 实操清单

按你 30 天计划 Day 16–19 共 8 小时(每天 2 小时):

Day 16(2h):Exercise 1 上半场

  • 看视频 0:00 – 13:01(引言 + 起步代码),接着看 13:01 – 约 50:00(Ex1 前半:交叉熵长链 → 线性层2 → tanh)
  • 做《Ex1_全链路反向_练习.ipynb》STEP 1 – 9
  • 暂停 + 纸笔:Karpathy 每次说 "now let's derive",先暂停自己推 2 分钟

Day 17(2h):Exercise 1 下半场 + BN 分步

  • 看视频 约 50:00 – 1:26:31(BN 分步反向 → 线性层1 → embedding;1:05:17 有 Bessel 校正插曲,讲 n−1 的来历)
  • 做《BN反向传播_练习.ipynb》七连击(分步部分)
  • 回到 Ex1 notebook 完成 STEP 10 – 12,全链路通关

Day 18(2h):Exercise 2 + Exercise 3 融合公式(推导最硬的一天)

  • 看视频 1:26:31 – 1:36:37(Ex2)→ 做《Ex2_交叉熵简化_练习.ipynb》
  • 看视频 1:36:37 – 1:50:02(Ex3,Karpathy 纸上推 BN 融合公式)→ 做 BN notebook 的"最终挑战"cell
  • 预期会卡 3 次以上,这是 99% 的人的正常体验;推导全文见《BN反向推导》讲义第三节

Day 19(2h):Exercise 4 拼起来训练

  • 看视频 1:50:02 – 结尾
  • 做《Ex4_完整训练_练习.ipynb》:单批热身验证 → 训练循环 → 校准 BN → 评估 + 采样
  • 完成收官自测,把这 4 天的公式整理成个人"backward 速查表"

十二、心理预期 + 给 1B 训练的预演

心理预期

你会卡住,会沮丧,会觉得"我是不是真的不适合搞这个"。这是每个人看 Part 4 的体验,包括我(Claude)也会觉得这一节难度陡增。

应对策略:

  • 允许自己第一遍只懂 60%。第二遍回到细节,第三遍才能复述
  • BatchNorm 反向卡住了,跳过它先去做 Exercise 4,把训练跑通。BN 反向可以之后再回头啃
  • 每天 2 小时不够时,允许 Day 16-19 变成 Day 16-22。这一节多花时间值得
  • 实在卡死了,可以跳过 Part 4 先去 Part 5/6,回头再补。但是别永远跳过

给 1B 训练的预演

你训练 1B 模型时会反复用到 Part 4 的能力:

1B 训练场景需要 Part 4 的什么能力
训练 NaN 调试能定位到哪一层、哪个操作梯度爆炸
读 Flash Attention 源码看懂里面的 backward 推导
实现 LoRA / QLoRA知道哪些梯度需要算、哪些可以跳过
用 gradient checkpointing 省显存理解"重新算 forward 来省 backward 显存"的原理
混合精度训练(fp16/bf16)知道哪些算子在低精度下数值不稳定
写自定义 CUDA kernel必须手动写 forward + backward

Karpathy 在视频结尾说:完成 Part 4 后,你就拥有了"赤手空拳调试任何梯度问题"的能力。这是从"会用神经网络"到"懂神经网络"的最关键一跃。