Part 4 全链路反向推导
〇、四天学习路线与文件对照
| 天 | 视频段落 | 练习 notebook | 本讲义章节 |
|---|---|---|---|
| Day 16 | 0:00 – 约 50:00(Ex1 上半场) | Ex1_全链路反向_练习.ipynb STEP 1–9 | 第一、二节 |
| Day 17 | 约 50:00 – 1:26:31(Ex1 下半场:BN 分步等) | BN反向传播_练习.ipynb 七连击 + Ex1 STEP 10–12 | BN 专题讲义 |
| Day 18 | 1:26:31 – 1:50:02(Ex2 + Ex3 融合) | Ex2_交叉熵简化_练习.ipynb + BN notebook 最终挑战 | 第三节 + BN 讲义第三节 |
| Day 19 | 1:50:02 – 结尾(Exercise 4) | Ex4_完整训练_练习.ipynb | 第四节 |
每个练习都有对应的 _参考答案.ipynb。Ex1 里 BN 一段可先"当黑盒"直接用给出的公式(走法 B),Day 18 再专攻。所有公式均已用数值梯度独立验证(误差 1e-10 量级)。
一、三条黄金法则(贯穿全部四天)
| 法则 | 内容 | 怎么用 |
|---|---|---|
| ① 形状大法 | 梯度形状永远等于原变量形状:dx.shape == x.shape |
写完公式先查形状。矩阵乘顺序记不住?用形状反推唯一合法的组合 |
| ② 广播必 sum | forward 里被 broadcast 复制的变量,backward 对被复制的维度求和 | b2 是 [27] 被复制 32 次 → db2 = dlogits.sum(0);counts_sum_inv 是 [32,1] 沿 dim=1 复制 → .sum(1, keepdim=True) |
| ③ 分叉必相加 | 一个变量在 forward 里被用了 k 次,backward 就有 k 条支路,梯度相加 | counts 既直接进 probs 又进 counts_sum;logits 既进 norm_logits 又进 logit_maxes |
二、Exercise 1:拆解版 cross-entropy 的逐步推导
平时一行 F.cross_entropy(logits, Yb),内部其实是 7 个算子:
logit_maxes = logits.max(1, keepdim=True).values # [32,1] 数值稳定
norm_logits = logits - logit_maxes # [32,27]
counts = norm_logits.exp() # [32,27]
counts_sum = counts.sum(1, keepdim=True) # [32,1]
counts_sum_inv = counts_sum ** -1 # [32,1]
probs = counts * counts_sum_inv # [32,27] = softmax
logprobs = probs.log() # [32,27]
loss = -logprobs[range(n), Yb].mean() # 标量
逐步 backward(倒序)
| 步 | Forward | Backward | 要点 |
|---|---|---|---|
| 1 | loss = -logprobs[range(n),Yb].mean() |
dlogprobs = zeros_like(logprobs)dlogprobs[range(n),Yb] = -1/n |
索引选中的元素才有梯度;mean → 1/n;负号别丢 |
| 2 | logprobs = probs.log() |
dprobs = (1/probs) * dlogprobs |
log 的导数用输入(probs),不能用输出 |
| 3 | probs = counts * counts_sum_inv |
dcounts_sum_inv = (counts*dprobs).sum(1,kd)dcounts_b1 = counts_sum_inv * dprobs |
广播在 dim=1 → sum(1);counts 有分叉,先存支路 1 |
| 4 | counts_sum_inv = counts_sum**-1 |
dcounts_sum = -counts_sum**-2 * dcounts_sum_inv |
幂函数 d(u⁻¹) = -u⁻² |
| 5 | counts_sum = counts.sum(1,kd) |
dcounts = dcounts_b1 + ones_like(counts)*dcounts_sum |
sum 反向 = 复制扩展;支路汇合相加 |
| 6 | counts = norm_logits.exp() |
dnorm_logits = counts * dcounts |
exp 的导数 = 输出自己(和 tanh 一样用输出省计算) |
| 7 | norm_logits = logits - logit_maxeslogit_maxes = logits.max(1).values |
dlogit_maxes = (-dnorm_logits).sum(1,kd)dlogits = dnorm_logits + one_hot(argmax)*dlogit_maxes |
max 反向:梯度只流向每行最大值位置;logits 分叉汇合 |
exp 上溢(e⁸⁰ 就 inf 了),对梯度没有贡献。Karpathy 专门停下来讲这件事:数值技巧不应改变数学。
exp(导数是自己)和 tanh(1−y²)用输出更省;log(1/x)、幂函数(a·u^(a−1))必须用输入。判断标准:导数能否用输出表达。写融合 kernel 时这决定你要缓存哪个张量——省显存的关键。
bias=False(GPT、LLaMA 都是)。手动 backward 让你亲眼看见这个结论,而不是背下来。
三、Exercise 2:从 7 步到 3 行的数学推导
对单个样本(先不管 batch 维),记 logits 为 l ∈ ℝ²⁷,正确类别 y,softmax 概率:
loss 展开(这一步是化简的关键——先取 log 再求导,而不是先算 p 再求导):
对任意 lj 求导,两项分别处理:第一项当 j=y 时贡献 −1 否则 0;第二项是复合函数:
合起来:
batch 上 loss 取 mean,再除以 n。翻译成 PyTorch 就是那 3 行:
dlogits = F.softmax(logits, 1) # p
dlogits[range(n), Yb] -= 1 # p − one_hot
dlogits /= n # mean 平摊
为什么它更稳、更快
- 稳:长链里的
dprobs = dlogprobs / probs在 probs → 0 时爆炸;简化版全程无除法(除以 n 除外),无 log(0) 风险。这就是F.cross_entropy接受 logits 而非 probs 的根本原因。 - 快:7 个中间张量的读写 → 1 次融合计算。这正是"算子融合"的原型 —— Flash Attention 对 softmax backward 做的是同一件事的高配版。
- 美:梯度就是"差距"本身。预测对了(py→1),梯度→0,不再更新;错得离谱,梯度接近满值。损失函数的设计哲学浓缩在这一行里。
验证性质:每行梯度之和 = Σjpj − 1 = 0。概率此消彼长,总量守恒 —— Ex2 notebook 里的可视化 cell 能亲眼看到这个"推与拉"结构。
四、Exercise 4:拼装与训练的五个要点
先热身再进循环:单 batch 上把拼装好的 backward 和 autograd cmp 一遍,7 个参数全 ✅ 再开始训练。循环里没有 autograd 兜底,错了不会报错,只会默默训歪。
整个循环包在
torch.no_grad()里:我们不需要 PyTorch 建计算图(这正是手动版反而更快的原因之一——省掉了建图和中间缓存开销)。forward 也不再拆 cross-entropy 长链,直接F.cross_entropy,因为反向用的是 Ex2 简化版。grads列表顺序必须与parameters严格对应:zip(parameters, grads)不会检查你有没有错位——形状恰好兼容时(broadcast)甚至不报错,直接训歪。这是分布式训练里"梯度桶错位"这类 bug 的微缩版。推理前用全训练集校准 BN 统计量:训练时 BN 用的是 batch 统计量,推理时单样本没有"batch"可言 → 用整个训练集的 mean/var 固定住。(真实 BN 层用 running average 边训边攒,视频里为了教学用一次性校准。)
结果预期:20000 步 loss 到 ~2.2 量级即拼装成功;跑满 200000 步 val ≈ 2.10,和 Part 3 autograd 版一致 —— 这就是通关铃声。🔔
五、调试检查清单(cmp ❌ 时按序排查)
- 形状:先 print 你的梯度和
t.grad的 shape。不等 → 八成是 keepdim 或 sum 维度错。 - sum 的维度:CE 链 dim=1,BN 链 dim=0,线性层 bias dim=0。对着 forward 里被广播的维度查。
- 分叉漏支路:counts(2 条)、logits(2 条)、bndiff(2 条)、hprebn(2 条)。maxdiff 不小不大(1e-3 ~ 1e-1)常是漏了一条支路。
- 输入 vs 输出:log/幂用输入,exp/tanh 用输出。用错了 maxdiff 会很大。
- notebook 重跑坑:cmp 莫名从 ✅ 变 ❌ → Restart & Run All。
- maxdiff 读数:1e-9 级 = 过关;1e-7 级 = 可接受的浮点误差;更大 = 公式错。
六、连回 1B 主线
| 这四天练的 | 1B 训练里对应的 |
|---|---|
| Ex1 逐算子 backward + 形状纪律 | 读/写自定义 CUDA kernel 的 backward;FSDP 梯度通信的形状契约 |
| Ex2 融合简化 | 所有 LLM 的 loss 实现;fused kernel 的设计思想 |
| Ex1 洞察 3(db1≈0) | 为什么 GPT/LLaMA 的 Linear 都 bias=False |
| Ex4 手动训练循环 | 训练 NaN 时你能逐层 print 梯度定位,而不是重启祈祷 |