Makemore Part 4 讲义:Backprop Ninja
〇、本节的位置与意义
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 1 | 13:01 – 1:26:31 | 手动 backprop 整个 forward 图(20+ 个变量,含分步 BN)——全片最长的一段 | ★★★ |
| 插曲 | 1:05:17 | Bessel 校正:方差为什么除 n−1(在 Ex1 中段出现) | ★★ |
| Exercise 2 | 1:26:31 – 1:36:37 | cross_entropy 的优雅简化(一行搞定) | ★★★★ |
| Exercise 3 | 1:36:37 – 1:50:02 | BatchNorm 的反向(视频最难的 20 分钟) | ★★★★★ |
| Exercise 4 | 1: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)✓
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
预期时长: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
但 mean 和 var 是"聚合操作",反向传播时梯度要"分回去给每个样本",公式变得超级复杂:
dhprebn = (bngain * bnvar_inv / n) * (
n * dhpreact
- dhpreact.sum(0)
- n/(n-1) * bnraw * (dhpreact * bnraw).sum(0)
)
5 个张量操作压缩在一行里。Karpathy 用 20 分钟在纸上推这一个公式。
关键洞察:BatchNorm 的反向之所以复杂,是因为 forward 里每个样本都参与了 mean 和 var 的计算,所以每个样本的梯度都依赖于 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 = doutdb = dout |
有 broadcast 时记得 .sum(0) |
y = a * b(element-wise) |
da = b * doutdb = a * dout |
注意是 element-wise,不是矩阵乘 |
y = x @ W |
dx = dout @ W.TdW = 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_entropycombined |
dlogits = softmax(logits)dlogits[range(n), y] -= 1dlogits /= n |
Exercise 2 的"优雅简化" |
BatchNormcombined |
见 Exercise 3,5 个操作压缩在一行 | Exercise 3 的"大魔王" |
六、值得记住的几个"坑"
梯度形状 vs 原变量形状必须严格一致
db.shape == b.shape,dW.shape == W.shape。形状不对 = 数值一定不对。有 broadcast 就要
.sum()
比如y = x + b,b 是 [D] 而 y 是 [N, D],反向时db = dout.sum(0),绝不能只是db = dout。+= 还是 =
正在初始化的 grad 用=,要"累加上一层传来的梯度"时用+=。Micrograd 时代说过的规则。tanh 的导数用 y 而非 x
公式dx = (1 - y**2) * dout,其中y = tanh(x)。用 y 省一次 tanh 计算(forward 时已经算过 y 了)。cross_entropy 的简化版要除以 N
公式dlogits /= n这一行**不能忘**。因为原始 loss 是.mean()而非.sum()。embedding 的 backward 必须用
index_add_不能用赋值
因为同一个字符可能在 batch 内出现多次,梯度要累加而非覆盖。BatchNorm 的
n / (n-1)因子很关键
PyTorch 的.var()默认unbiased=True,分母是n-1而非n。反向公式要带这个因子。验证用
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 篇
- 线性层
y = x @ W + b的 backward 三个公式分别是? - 为什么
db = dout.sum(0)不是db = dout? - tanh 的 backward 用 y 还是 x 算?为什么用 y 更聪明?
view操作的 backward 是什么?需不需要算?
Exercise 2 篇
- cross_entropy 的简化公式
(probs - one_hot) / N是怎么推出来的? - 为什么这个简化在数值上比 "softmax → log → NLL" 更稳定?
- 这跟
F.cross_entropy(logits, y)内部的实现一样吗?
Exercise 3 篇
- BatchNorm 反向为什么比一般算子复杂这么多?
- BN 公式里
n/(n-1)这个因子从哪来的? - 如果 batch_size = 1,BN backward 会出什么问题?
Exercise 4 篇
- 手动 backward 训练出的模型,val loss 应该跟 Part 3 完全一致还是仅近似?
- 为什么手动 backward 比
loss.backward()慢得多?
十、看完后的自我检测
这 10 道题答不上来 3 道以上,回去重看:
| # | 题目 | 标准答案要点 |
|---|---|---|
| 1 | 写出 y = x @ W + b 的三个 backward 公式 | dx = dout @ W.T, dW = x.T @ dout, db = dout.sum(0) |
| 2 | tanh backward 的公式 | dx = (1 - y**2) * dout,y 是 tanh 输出 |
| 3 | cross_entropy 简化 backward 三行代码 | dlogits = softmax(logits); dlogits[range(n),y] -= 1; dlogits /= n |
| 4 | broadcasting 加法的 backward 为什么要 sum | 因为 b 被复制 N 次,每次都参与 forward |
| 5 | embedding 反向用 index_add_ 不用赋值的原因 | 同一字符可能出现多次,梯度要累加 |
| 6 | view 的 backward 是什么 | dx = dout.view(x.shape),不算数 |
| 7 | BatchNorm 反向用 mean 的部分对应公式哪一项 | -dhpreact.sum(0) 这一项 |
| 8 | cmp() 验证的 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 训练的预演
心理预期
应对策略:
- 允许自己第一遍只懂 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 后,你就拥有了"赤手空拳调试任何梯度问题"的能力。这是从"会用神经网络"到"懂神经网络"的最关键一跃。