跳到主要内容

自动求导与计算图 — 框架到底替你做了什么

这一章讲三件事: 深度学习为什么离不开自动求导;计算图是什么、反向传播怎么沿它走; 以及同一个回归任务用三种方式(NumPy 手工、autograd、TensorFlow 静态图)实现后的对照—— 看完你就知道「框架」这个词背后,真正被包起来的是哪几件事。

1. 顶层全景:正向记路,反向照路走

前向传播(算出结果,顺便记下每步) 反向传播(照记录往回推)
x ──┬─ mul → y ─┬─ add → z ── loss loss
w ──┘ │ ↑ 链式法则逐层回代
b ──────────────┘ y.grad ← z.grad ← ...

w.grad、b.grad 落进各自的 .grad

图说:图上的圆是变量(数据),方是算子(运算)。正向只算一遍数值;
反向不重算数值,只把「损失对每个变量的导数」沿刚才的记录传回来。

反向传播(backpropagation)本身不是深度学习发明的——它就是微积分的链式法则(复合函数的导数等于各层导数相乘)套在图上逐层执行。框架的贡献是:你只写正向,反向它照着正向的记录自动生成。这套机制在 PyTorch 里叫 autograd1

2. 计算图:z = wx + b 长什么样

先看最小的例子。表达式 z = wx + b 在框架眼里不是一行代码,而是三个节点、两个算子的图:

x(叶子) ──┐
mul → y(中间)
w(叶子) ──┘ │
add → z(根)
b(叶子) ────────────┘

三类节点各有名字,也各有固定行为:叶子节点是用户直接创建的变量(x、w、b);中间节点是算出来的(y、z);最末端要优化的叫根。谁需要求导,创建时标 requires_grad=True——默认 False,凡依赖它的后续节点自动变 True2。每个算出来的节点身上还挂着 grad_fn,记录「我是被哪个算子算出来的」;叶子节点没有,为 None3

对根调用 z.backward(),autograd 从 z 出发沿图反向溯源,把每个叶子的梯度累加进它的 .grad 属性4

书里把这句跑成了数:x 取 2,w、b 随机、都开 requires_grad,backward 之后打印——w.grad 是 2.0(z 对 w 的导数就是 x 的值),b.grad 是 1.0(z 对 b 的导数恒为 1),而 x 无须求导、梯度为 None5两个数字各就各位,这就是「自动求导」的全部观感。

两条容易踩的规矩,都源于「梯度是累加的」这一设计:非叶子节点(y、z)的梯度在 backward 之后立即清空,不看就没了;中间缓存(算好的中间结果)也一并清掉,想再 backward 一次要传 retain_graph=True,且第二次 backward 的结果会在第一次上,通常要先手动清零6。评估、测试阶段不想要任何梯度记录,用 with torch.no_grad() 包住代码块即可7

3. 主走查:同一个回归任务的三种写法

本章的主走查是一道回归题,书里用三种方式各写了一遍,这三种写法就是「框架价值」的对照实验。

任务: 造一组数据点,真实规律是 y = 3x² + 2,外加一点噪声;让模型只拿到 x 和 y,自己学出 w=3、b=2(模型假设 y = wx² + b)8

写法一:纯 NumPy,梯度手推。 损失取「预测值减真值」的平方和的一半,对 w、b 求偏导的式子自己写,自己代入,再按「参数 -= 学习率 × 梯度」更新。循环 800 遍后,w 收敛(误差不再明显变化)到 2.9586、b 收敛到 2.1018——离真值 3 和 2 很近9

写法二:Tensor + autograd,梯度免手推。 数据换成 Tensor,w、b 标 requires_grad=True;循环里只剩四步:算预测 → 算损失 → loss.backward() → 手动更新参数(更新包在 no_grad 里,且每一遍手动把梯度清零,因为梯度是累加的)10。同样 800 遍,w=2.9645、b=2.1146,与写法一几乎重合11对比看:损失到梯度的那一段推导,从十几行数学变成了一个函数调用(一行代码的事)。

写法三:TensorFlow 静态图。 先用占位符声明「数据之后从这进」,把整个计算(包括参数更新)先搭成一张固定的图,再开 Session 喂数据执行。跑 2000 轮(把数据完整过 2000 遍),损失降到 0.0038,w≈2.92、b≈2.1212

三种写法摆在一起,差别收敛成两句话:

梯度从哪来参数更新在哪
NumPy 手工自己推导、自己写式子普通赋值语句
PyTorch autogradloss.backward() 自动图外的普通语句
TensorFlow 静态图tf.gradients 声明成图节点图内,更新本身是图的一部分13

判断(我们的,不是书里的): 「更新在图内还是图外」这行差别,是静态图时代与动态图时代真正的分水岭——图外更新意味着训练循环就是普通 Python 代码,可以随时打断点、print、换逻辑;这也解释了为什么科研圈倒向 PyTorch。今天两个框架都收敛到动态执行(见 §5),这一行差别已进入历史,但读 2020 年前的代码和论文时它是必备背景。 如果错,会错在: 如果性能敏感的生产部署仍大规模依赖先建图后优化(今天的推理——部署后跑预测——引擎确实这么做),那「图外更新」的优势就只属于调试期,不属于全生命周期——我们的表述就只对了一半。

4. 非标量反向传播:雅可比的一小步

上一节的损失是标量(一个数),backward 不用传参数。如果要对一个向量求导,PyTorch 干脆禁止——张量对张量的导数是个矩阵(所谓雅可比矩阵,Jacobian:输出的每个分量对输入的每个分量分别求导,排成的表),框架规定只允许「标量对张量」,要反传非标量就得先把它变回标量问题14

变法是乘一个形状相同的权重向量 v,把「向量对向量的导数」变成「vᵀ·loss 这个标量对输入的导数」——也就是拿雅可比矩阵跟 v 做一次矩阵乘。书里用一组能手算的式子验证:y1 = x1² + 3x2,y2 = x2² + 2x1,取 x=(2,3)。手工展开可得雅可比 J = [[4,3],[2,6]] (每格是某个 y 对某个 x 的偏导)。

直接调 y.backward(torch.Tensor([[1,1]])),得 x.grad = [6, 9]——这不是雅可比,而是 J 的两行之和(权重 [1,1] 把两行加在了一起)15。想拿到完整 J 要分两步:先传 v=(1,0) 拿第一行,retain_graph=True 再传 v=(0,1) 拿第二行,中间手动把 x.grad 清零,最后得到 [[4,3],[2,6]],与手算一致16

这段的价值不在技巧,在认知:深度学习里几乎所有的「求导」实际都是「v·J」这种向量-矩阵乘——损失是标量时 v=1,仅此而已。框架的禁令不是限制,是把数学事实摆上台面。

一个能立刻校准直觉的旁证:有人用 94 行 Python 从零写了一个叫 micrograd 的标量级自动求导引擎,backward 就是把计算图按拓扑序反向走一遍、逐节点套链式法则(依据: shelf=ai-frontier-reference/micrograd@src:micrograd/engine.py:54,事实=整个引擎仅 94 行,backward 以拓扑序遍历图)。自动求导不是魔法,是把链式法则代码化。

5. 动态图赢在哪:调试与「随时改」

第 3 节的三种写法里还藏着一个话题:为什么书把动态图当卖点。动态图(dynamic computational graph)指计算图在每次前向传播(正向算一遍)时当场构建、用完即弃,下一次可以完全不同。

TensorFlow 1.x 那类静态图则要先定义好整张图,再开一个会话(执行环境)喂数据执行17

动态、且采用基于磁带(tape)的记录方式,意味着每执行一条命令立刻能看结果,出错当场断点排查,训练中途改网络结构也不影响性能18。书里点了一句后来被验证的判断:TensorFlow 2.0 新增的 Eager Execution(即时执行模式)已把它自己的默认执行方式转向动态——两大框架在此汇合19

判断(我们的,不是书里的): 今天读这本书,「动态 vs 静态」之争本身已经是历史注脚,但书里给的两条判断(动态利于实验、科研圈因此偏爱 PyTorch)对理解 2017-2020 年的生态迁移依然是最省力的解释框架。 如果错,会错在: 如果静态图在高性能部署场景始终保有独立生态(后来的推理优化器确实部分如此),那「汇合」就只是训练侧的汇合,不是全链路的。

6. 作者的判断与证据

给了证据的: 三种实现对同一任务的结果数字(2.9586/2.1018、2.9645/2.1146、0.0038)都是书里实际运行的输出,可复算;标量与非标量反向传播都配了手算对照,autograd 输出与手工雅可比逐格核对。

给了理由、但属作者立场: 「PyTorch 动态图在调试方面非常方便,通过断点检查就可以高效解决问题」出自前言,这是作者选 PyTorch 的第二条理由(「前言」第 19 段,text/01-fm.txt:20,搜「断点检查」)。理由真实,但属体验判断而非实验结论。

书里没展开的: backward 内部如何分配内存、如何调度算子,书完全没讲——对读者合适,但要知道「计算图执行引擎」这一层被刻意略过了。

7. 边界与局限

  • 梯度清零是手工活。 书里两处训练循环都要手动 zero_grad,漏写不会报错、只会让梯度无限累加——这是初学者最常见的静默 bug。第 3 章的 optimizer.zero_grad() 会把它规范化。
  • 本书不讲二阶导与高阶图。 create_graph 参数一笔带过,实际用在元学习等场景,书未覆盖。
  • TensorFlow 部分的代码今天已双重过时: 既是静态图习语,又用了 tf.placeholder 这类 2.0 已移除的 API。读它当「历史对照」,不要照抄。

8. 可带走的

  1. 前向传播「顺便」记下每步运算,这张记录就是计算图;反向传播=沿图反向套链式法则,autograd 全自动;
  2. 叶子节点(requires_grad=True)存梯度,中间节点的梯度用完即弃;记住「梯度累加」,所以每一遍训练前要清零;
  3. backward 默认只对标量;对向量求导=传一个同形状的 gradient 参数,拿到的永远是「雅可比×这个参数」,不是雅可比本身;
  4. PyTorch 参数更新在图外(普通 Python 语句),静态图框架在图内——这决定了谁能用普通调试器;
  5. no_grad 是评估阶段的标准姿势,省内存也省时间;
  6. 想知道框架替你做了什么,就把同一个任务用 NumPy 手写一遍——梯度那十几行推导就是被包掉的全部。

9. 原文地图

主题原书章原文位置
反向模式自动微分、动态图第2章 PyTorch基础text/03-ch02-2-pytorch.txt:31(搜「动态计算图」) · text/03-ch02-2-pytorch.txt:34(搜「Reverse-mode」)
requires_grad 与依赖传播第2章 PyTorch基础text/03-ch02-2-pytorch.txt:551(搜「requires_grad」)
grad_fn 属性第2章 PyTorch基础text/03-ch02-2-pytorch.txt:559(搜「grad_fn」)
梯度累加、retain_graph第2章 PyTorch基础text/03-ch02-2-pytorch.txt:570(搜「retain_graph」)
no_grad 用于测试第2章 PyTorch基础text/03-ch02-2-pytorch.txt:574(搜「no_grad」)
动态图每次前向重建第2章 PyTorch基础text/03-ch02-2-pytorch.txt:578(搜「重新构建」)
计算图与叶子节点第2章 PyTorch基础text/03-ch02-2-pytorch.txt:583(搜「有向无环」) · text/03-ch02-2-pytorch.txt:551(搜「叶子节点」)
链式法则反向溯源第2章 PyTorch基础text/03-ch02-2-pytorch.txt:593(搜「链式法则」)
标量反传实例(x=2)第2章 PyTorch基础text/03-ch02-2-pytorch.txt:613(搜「torch.Tensor([2])」) · text/03-ch02-2-pytorch.txt:650(搜「梯度分别为」)
禁止张量对张量求导第2章 PyTorch基础text/03-ch02-2-pytorch.txt:661(搜「不让张量」)
gradient 参数与雅可比第2章 PyTorch基础text/03-ch02-2-pytorch.txt:666(搜「标量对张量」) · text/03-ch02-2-pytorch.txt:669(搜「雅可比」)
[1,1] 得两行之和第2章 PyTorch基础text/03-ch02-2-pytorch.txt:708(搜「6., 9.」)
两步取完整雅可比第2章 PyTorch基础text/03-ch02-2-pytorch.txt:570(搜「retain_graph」) · text/03-ch02-2-pytorch.txt:729(搜「4., 3.」)
回归任务与真实规律第2章 PyTorch基础text/03-ch02-2-pytorch.txt:742(搜「3x2」)
NumPy 手工梯度结果第2章 PyTorch基础text/03-ch02-2-pytorch.txt:798(搜「grad_w」) · text/03-ch02-2-pytorch.txt:816(搜「2.95859544」)
autograd 版与手动更新第2章 PyTorch基础text/03-ch02-2-pytorch.txt:864(搜「loss.backward」) · text/03-ch02-2-pytorch.txt:866(搜「no_grad」)
autograd 版结果第2章 PyTorch基础text/03-ch02-2-pytorch.txt:887(搜「2.9645」)
TF Eager Execution第2章 PyTorch基础text/03-ch02-2-pytorch.txt:896(搜「Eager Execution」)
TF 更新在图内第2章 PyTorch基础text/03-ch02-2-pytorch.txt:936(搜「计算图的一部分内容」)
TF 结果 0.0038第2章 PyTorch基础text/03-ch02-2-pytorch.txt:976(搜「0.0038」)
动态图交互式第2章 PyTorch基础text/03-ch02-2-pytorch.txt:982(搜「交互式」)

Footnotes

  1. 出处:「第2章 PyTorch基础」第 542 段(text/03-ch02-2-pytorch.txt:542,搜「自动求导」)。

  2. 出处:「第2章 PyTorch基础」第 551 段(text/03-ch02-2-pytorch.txt:551,搜「requires_grad」)。原文:缺省 False,设 True 后「与之有依赖关系的节点会自动变为True」。

  3. 出处:「第2章 PyTorch基础」第 559 段(text/03-ch02-2-pytorch.txt:559,搜「grad_fn」)。

  4. 出处:「第2章 PyTorch基础」第 562 段(text/03-ch02-2-pytorch.txt:562,搜「累加」)与第 593 段(text/03-ch02-2-pytorch.txt:593,搜「链式法则」)。

  5. 出处:「第2章 PyTorch基础」第 613 段(text/03-ch02-2-pytorch.txt:613,搜「torch.Tensor([2])」)与第 650 段(text/03-ch02-2-pytorch.txt:650,搜「梯度分别为」)。w.grad=2、b.grad=1、x.grad=None 均为原文运行结果。

  6. 出处:「第2章 PyTorch基础」第 570 段(text/03-ch02-2-pytorch.txt:570,搜「retain_graph」)与第 653 段(text/03-ch02-2-pytorch.txt:653,搜「自动清空」)。

  7. 出处:「第2章 PyTorch基础」第 574 段(text/03-ch02-2-pytorch.txt:574,搜「no_grad」)。

  8. 出处:「第2章 PyTorch基础」第 742 段(text/03-ch02-2-pytorch.txt:742,搜「3x2」)。

  9. 出处:「第2章 PyTorch基础」第 798 段(text/03-ch02-2-pytorch.txt:798,搜「grad_w」)与第 816 段(text/03-ch02-2-pytorch.txt:816,搜「2.95859544」)。迭代次数 800 见第 791 段(text/03-ch02-2-pytorch.txt:791,搜「range(800)」)。

  10. 出处:「第2章 PyTorch基础」第 864 段(text/03-ch02-2-pytorch.txt:864,搜「loss.backward」)与第 866 段(text/03-ch02-2-pytorch.txt:866,搜「no_grad」);梯度清零见第 871 段(text/03-ch02-2-pytorch.txt:871,搜「梯度清零」)。

  11. 出处:「第2章 PyTorch基础」第 887 段(text/03-ch02-2-pytorch.txt:887,搜「2.9645」)。

  12. 出处:「第2章 PyTorch基础」第 976 段(text/03-ch02-2-pytorch.txt:976,搜「0.0038」)。

  13. 出处:「第2章 PyTorch基础」第 936 段(text/03-ch02-2-pytorch.txt:936,搜「计算图的一部分内容」)。原文明确对比:「而PyTorch,这部分属于计算图之外」。

  14. 出处:「第2章 PyTorch基础」第 661 段(text/03-ch02-2-pytorch.txt:661,搜「不让张量」)与第 666 段(text/03-ch02-2-pytorch.txt:666,搜「标量对张量」)。

  15. 出处:「第2章 PyTorch基础」第 708 段(text/03-ch02-2-pytorch.txt:708,搜「6., 9.」)。原文自己指出「这个结果与我们手工运算的不符」并分析原因在 v 的取值。

  16. 出处:「第2章 PyTorch基础」第 570 段(text/03-ch02-2-pytorch.txt:570,搜「retain_graph」)与第 729 段(text/03-ch02-2-pytorch.txt:729,搜「4., 3.」)。

  17. 出处:「第2章 PyTorch基础」第 578 段(text/03-ch02-2-pytorch.txt:578,搜「重新构建」)。

  18. 出处:「第2章 PyTorch基础」第 31 段(text/03-ch02-2-pytorch.txt:31,搜「动态计算图」)与第 982 段(text/03-ch02-2-pytorch.txt:982,搜「交互式」)。

  19. 出处:「第2章 PyTorch基础」第 896 段(text/03-ch02-2-pytorch.txt:896,搜「Eager Execution」)。