Stephen 技术博客

AI和数据平台的工程实践

← 返回文章列表

Part 4 全链路反向推导

覆盖 Exercise 1(全链路)· Exercise 2(交叉熵简化)· Exercise 4(完整训练)
Exercise 3(BatchNorm)另见专题讲义《Makemore_Part4_讲义_BN反向推导.html》
Day 16 – 19 全套

〇、四天学习路线与文件对照

视频段落练习 notebook本讲义章节
Day 160:00 – 约 50:00(Ex1 上半场)Ex1_全链路反向_练习.ipynb STEP 1–9第一、二节
Day 17约 50:00 – 1:26:31(Ex1 下半场:BN 分步等)BN反向传播_练习.ipynb 七连击 + Ex1 STEP 10–12BN 专题讲义
Day 181:26:31 – 1:50:02(Ex2 + Ex3 融合)Ex2_交叉熵简化_练习.ipynb + BN notebook 最终挑战第三节 + BN 讲义第三节
Day 191: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
dim=0 还是 dim=1? Part 4 里最阴险的坑。判断方法只有一个:看 forward 里哪个维度被广播/聚合。cross-entropy 链的聚合都在 dim=1(每行是一个 27 类的概率分布:counts_sum、logit_maxes 都是 [32,1]);BN 的聚合都在 dim=0(对 batch 内 32 个样本求统计量:bnmeani、bnvar 都是 [1,64])。抄公式不看维度,必翻车。

二、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(倒序)

ForwardBackward要点
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_maxes
logit_maxes = logits.max(1).values
dlogit_maxes = (-dnorm_logits).sum(1,kd)
dlogits = dnorm_logits + one_hot(argmax)*dlogit_maxes
max 反向:梯度只流向每行最大值位置;logits 分叉汇合
洞察 1:dlogit_maxes ≈ 0(量级 1e-17)。softmax 对每行减去任意常数不变(分子分母同乘 e⁻ᶜ 抵消),所以 loss 对 logit_maxes 根本不敏感——减 max 纯粹是为了防止 exp 上溢(e⁸⁰ 就 inf 了),对梯度没有贡献。Karpathy 专门停下来讲这件事:数值技巧不应改变数学
洞察 2:导数用输入还是输出? 一个省心的规律:exp(导数是自己)和 tanh(1−y²)用输出更省;log(1/x)、幂函数(a·u^(a−1))必须用输入。判断标准:导数能否用输出表达。写融合 kernel 时这决定你要缓存哪个张量——省显存的关键。
洞察 3:db1 ≈ 0。走到 STEP 11 你会发现 b1 的梯度小得反常。因为 BN 第一步就是减均值:b1 加上去多少,bnmeani 就减掉多少 —— b1 是个"无用参数"。这就是为什么带 Norm 层的 Linear 通常设 bias=False(GPT、LLaMA 都是)。手动 backward 让你亲眼看见这个结论,而不是背下来。

三、Exercise 2:从 7 步到 3 行的数学推导

对单个样本(先不管 batch 维),记 logits 为 l ∈ ℝ²⁷,正确类别 y,softmax 概率:

pj = elj / Σk elk

loss 展开(这一步是化简的关键——先取 log 再求导,而不是先算 p 再求导):

L = −log py = −ly + log Σk elk

对任意 lj 求导,两项分别处理:第一项当 j=y 时贡献 −1 否则 0;第二项是复合函数:

∂/∂lj [log Σk elk] = elj / Σk elk = pj

合起来:

∂L/∂lj = pj − 𝟙{j = y}  (梯度 = 预测概率 − one-hot 真值)

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:拼装与训练的五个要点

  1. 先热身再进循环:单 batch 上把拼装好的 backward 和 autograd cmp 一遍,7 个参数全 ✅ 再开始训练。循环里没有 autograd 兜底,错了不会报错,只会默默训歪。

  2. 整个循环包在 torch.no_grad():我们不需要 PyTorch 建计算图(这正是手动版反而更快的原因之一——省掉了建图和中间缓存开销)。forward 也不再拆 cross-entropy 长链,直接 F.cross_entropy,因为反向用的是 Ex2 简化版。

  3. grads 列表顺序必须与 parameters 严格对应zip(parameters, grads) 不会检查你有没有错位——形状恰好兼容时(broadcast)甚至不报错,直接训歪。这是分布式训练里"梯度桶错位"这类 bug 的微缩版。

  4. 推理前用全训练集校准 BN 统计量:训练时 BN 用的是 batch 统计量,推理时单样本没有"batch"可言 → 用整个训练集的 mean/var 固定住。(真实 BN 层用 running average 边训边攒,视频里为了教学用一次性校准。)

  5. 结果预期:20000 步 loss 到 ~2.2 量级即拼装成功;跑满 200000 步 val ≈ 2.10,和 Part 3 autograd 版一致 —— 这就是通关铃声。🔔

五、调试检查清单(cmp ❌ 时按序排查)

  1. 形状:先 print 你的梯度和 t.grad 的 shape。不等 → 八成是 keepdim 或 sum 维度错。
  2. sum 的维度:CE 链 dim=1,BN 链 dim=0,线性层 bias dim=0。对着 forward 里被广播的维度查。
  3. 分叉漏支路:counts(2 条)、logits(2 条)、bndiff(2 条)、hprebn(2 条)。maxdiff 不小不大(1e-3 ~ 1e-1)常是漏了一条支路。
  4. 输入 vs 输出:log/幂用输入,exp/tanh 用输出。用错了 maxdiff 会很大。
  5. notebook 重跑坑:cmp 莫名从 ✅ 变 ❌ → Restart & Run All。
  6. 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 梯度定位,而不是重启祈祷