计算图在深度学习中有什么作用?
简化版
计算图把张量运算表示为有向无环图:前向按依赖计算值,反向依据链式法则从损失向叶参数传播梯度。动态图运行时建图灵活易调试,静态图先捕获再优化,通常更利于算子融合和部署。
详细版
-
节点代表张量或算子,边代表数据依赖与梯度流。
-
反向传播保存中间激活或用重计算换显存。
-
叶张量是否 requires_grad、detach 和原地修改会影响图连接。
-
同一节点收到多条下游路径的梯度时需要累加。
-
训练后应释放旧图,重复 backward 通常需显式 retain_graph。
完整版教学
一、计算图是链式法则的执行计划
复合函数的梯度需要每个局部算子的导数与前向中间值。
框架在前向记录父节点和 backward 函数,反向再按拓扑逆序执行。
分支图中参数可能通过多条路径影响损失,微积分要求这些路径贡献相加;这也是梯度默认累积而非覆盖的原因。
二、底层机制与公式
z = x*w
y = z*z + z
dy/dz = 2z + 1
dy/dw = (2z+1)*x
三、带数字的推演
令 x=2、w=3,则 z=6、y=42,dy/dw=(2×6+1)×2=26。
图会保存 z=6 供反向使用,而不是从最终 y 猜出局部导数。
四、方案对比
| 方案/对象 | 核心特点 | 代价或边界 |
|---|---|---|
| 动态图 | 边运行边建图 | 灵活,编译优化较难 |
| 静态/捕获图 | 先构建整体计划 | 融合强,控制流受约束 |
| 梯度检查点 | 丢弃部分激活再重算 | 省显存、增计算 |
五、执行流程
输入 -> 前向算子节点 -> loss
loss.grad=1 -> 逆拓扑调用 backward -> 路径梯度累加 -> 叶参数 grad
六、边界条件与工程代价
detach() 返回与原图断开的视图,适合停止梯度;误用会让前面参数永远收不到更新。
原地操作可能覆盖反向需要的版本,框架因此做 version counter 检查。
高阶梯度要求反向过程本身也被记录成图,通常需 create_graph=True;这会显著增加内存。
记忆钩子:把图记成“前向存依赖,反向走逆序,分支做累加”。
七、常见误区与追问
-
误区:计算图只在反向时创建。 动态图通常在前向执行算子时记录依赖。
-
追问:分支梯度为何相加? 损失对同一变量的总导数是所有路径贡献之和。
-
误区:zero_grad 是清空计算图。 它清的是参数梯度缓冲,图生命周期是另一件事。
-
追问:detach 与 no_grad 有何区别? 前者断开特定张量,后者在作用域内不记录运算。
-
追问:为何不能随意原地修改激活? 反向可能依赖修改前的值,版本不一致会报错或算错。
八、加强记忆
把图记成“前向存依赖,反向走逆序,分支做累加”。
理解节点、激活生命周期和断图操作后,autograd 的多数报错都能沿图定位。