Skip to content

计算图

计算图 (Computational graphs) 是许多深度学习框架实现和优化计算的基本概念,特别是用于梯度计算的反向传播 (backpropagation) 的关键过程。

计算图是一种有向图 (directed graph),其中:

  • 节点 (Nodes): 表示数学运算(例如,加法、乘法、激活函数 (activation functions))或变量(输入数据 (input data)、参数 (parameters) 如权重和偏置、中间值 (intermediate values))。
  • 边 (Edges): 表示数据流(张量 (tensors))在运算之间的流动。一个运算的输出成为另一个运算的输入。

本质上,它们提供了一种可视化和结构化表示复杂数学表达式的方式,将其分解为一系列简单的操作。

例如,考虑表达式:g = (x + y) * z

我们可以将其表示为图:

  1. x、y 和 z 的输入节点。
  2. 一个加法节点 (+),将 x 和 y 作为输入,产生中间结果 p = x + y。
  3. 一个乘法节点 (*),将 p 和 z 作为输入,产生最终结果 g。

边将指示从 x 和 y 流向 + 节点,从 z 流向 * 节点,以及从 + 节点(输出 p)流向 * 节点。

神经网络的前向传播 (forward pass) 自然可以表示为计算图。每一层都涉及矩阵乘法、加法(偏置)和激活函数,所有这些都成为图中的节点。输入数据和权重/偏置也是节点。

前向传播涉及通过输入值并按拓扑顺序 (topological order)(从输入到最终输出)计算每个节点的输出来评估图。

使用我们的示例 g = (x + y) * z,如果 x = 1,y = 3 并且 z = -3:

  1. + 节点计算 p = 1 + 3 = 4。
  2. * 节点计算 g = p * z = 4 * (-3) = -12。

这个过程反映了神经网络计算其预测的方式。

后向传播(通过自动微分 (Autograd))

Section titled “后向传播(通过自动微分 (Autograd))”

计算图在深度学习中的真正强大之处在于它们能够促进自动微分 (automatic differentiation) (autograd),以便进行后向传播 (backward pass)(反向传播 (backpropagation))。

后向传播的目标是计算最终输出(通常是损失函数 (loss function))相对于每个输入变量和参数(权重、偏置)的梯度 (gradient)。这些梯度是训练期间使用梯度下降更新参数所需的。

利用计算图结构和微积分的链式法则 (chain rule of calculus),框架可以通过将导数 (derivatives) 向后传播通过图来自动计算这些梯度。

对于我们的示例 g = (x + y) * z,令 p = x + y。我们想找到 ∂g/∂x、∂g/∂y 和 ∂g/∂z。

  1. 从输出开始: g 相对于自身的梯度 ∂g/∂g = 1。
  2. 通过 * 节点后向传播: 我们需要局部梯度 (local gradients) ∂g/∂p 和 ∂g/∂z。因为 g = p * z,我们知道 ∂g/∂p = z 和 ∂g/∂z = p。使用前向传播中的值 (p=4, z=-3),我们得到 ∂g/∂p = -3 和 ∂g/∂z = 4。
  3. 通过 + 节点后向传播: 我们需要 ∂p/∂x 和 ∂p/∂y。因为 p = x + y,我们有 ∂p/∂x = 1 和 ∂p/∂y = 1。
  4. 应用链式法则: 为了获得 g 相对于原始输入 x 和 y 的梯度,我们使用链式法则:
  5. ∂g/∂x = (∂g/∂p) * (∂p/∂x) = (-3) * 1 = -3
  6. ∂g/∂y = (∂g/∂p) * (∂p/∂y) = (-3) * 1 = -3

我们现在得到了所有期望的梯度:∂g/∂x = -3,∂g/∂y = -3,∂g/∂z = 4。

深度学习框架会自动执行此过程。通过定义前向传播(显式构建图或隐式通过运算),框架会记录运算序列,然后通过向后应用链式法则来计算梯度。

静态图 (Static Graphs)(例如,TensorFlow 1.x、Theano): 您首先定义整个计算图结构。然后,您可以多次执行此图的部分,输入不同的数据。这允许进行重要的提前优化 (ahead-of-time optimization),但对于调试和处理动态控制流 (dynamic control flow) 可能不够直观。

动态图 (Dynamic Graphs)(例如,PyTorch、TensorFlow 2.x Eager Execution): 图在 Python 中执行运算时动态构建。这感觉更像标准的编程,简化了调试,并且容易处理条件逻辑 (conditional logic) 或可变序列长度 (variable sequence lengths)。虽然最初可能优化程度较低,但像图编译 (graph compilation)(例如,TF2 中的 tf.function,PyTorch 中的 torch.compile)等技术可以捕获和优化动态执行的部分以提高性能。

理解计算图的概念有助于阐明自动微分的工作原理以及为什么它对于高效训练深度神经网络如此核心。