反向传播是一种算法,用于计算模型输出(通常是标量损失函数)相对于模型参数的变化率。它广泛应用于人工神经网络,将导数信息沿生成输出的运算过程反向传递。反向传播是反向模式自动微分的一种具体应用,而不是参数更新规则:优化器利用它计算出的梯度来调整模型。这一区别将变化方向的计算与如何沿该方向调整参数的决策分开。(jmlr.org)
数学原理
[ \frac{\partial L}{\partial u}
\frac{\partial L}{\partial v} \frac{\partial v}{\partial u}. ]
当 (u) 影响多个下游量时,需要将所有路径上的贡献相加。计算图用来表示这些依赖关系:节点表示数值或运算,边则表示运算使用哪些数值作为输入。反向传播按照依赖关系的逆序遍历计算图,将局部导数与已经算出的下游敏感度结合起来。因此,它不必沿每一条可能的路径,分别重新计算每个参数的贡献。(pytorch.org)
对于标量损失,反向计算从 (\partial L/\partial L=1) 开始。每个运算将传入的敏感度转换为其各个输入的敏感度。对于向量值中间量,这相当于乘以局部雅可比矩阵的转置,通常无需显式构造完整矩阵。同样的机制通过累加各项贡献,也能处理分支计算和共享参数。(docs.pytorch.org)
前向与反向计算
考虑一个输入为 (a^0=x) 的网络。在第 (l) 层,权重矩阵 (W^l)、偏置向量 (b^l) 和逐元素作用的激活函数 (\phi_l) 产生
[ z^l=W^l a^{l-1}+b^l, \qquad a^l=\phi_l(z^l). ]
前向计算先求出这些表达式的值,再计算损失。求导所需的数值会被保留下来,供反向计算使用。将该层的敏感度定义为 (\delta^l=\partial L/\partial z^l),则隐藏层的递推关系为
[ \delta^l= \left((W^{l+1})^\top\delta^{l+1}\right) \odot\phi_l'(z^l), ]
其中 (\odot) 表示逐元素相乘。参数的导数为
[ \frac{\partial L}{\partial W^l}
\delta^l(a^{l-1})^\top, \qquad \frac{\partial L}{\partial b^l}=\delta^l. ]
这些表达式描述的是单个样本的情形;对于一个批次,导数会按照损失函数采用求和还是求均值的约定进行合并。更一般的网络架构会对各个运算应用同样的链式求导过程,而不是依赖这一特定的逐层递推关系。(deeplearningbook.org)
与学习和优化的关系
反向传播为数学优化提供梯度。采用梯度下降时,参数 (\theta) 按以下规则更新:
[ \theta\leftarrow\theta-\eta\nabla_\theta L, ]
其中 (\eta) 为学习率。随机梯度下降使用抽样得到的样本或小批量样本来估计梯度。其他优化器也可以使用相同的导数,但采用不同的更新规则。因此,一次训练迭代可分为前向求值、梯度计算和参数更新三个环节。(deeplearningbook.org)
在监督学习中,目标函数通常将预测结果与训练数据中的目标值进行比较。隐藏层没有单独的目标值;其梯度由它们对最终损失的贡献推导而来。这使得表征学习成为可能,即内部特征通过训练形成,而非完全由人事先指定。戴维·鲁梅尔哈特、杰弗里·辛顿和罗纳德·威廉斯在1986年的实验中展示了多层网络的这一特性。(nature.com)
效率与实现
对于由标准可微运算组成的计算,求取标量输出的完整梯度所需的算术运算量,通常仅为计算输出本身的常数倍。这使得反向模式适合参数众多、输出相对较少的模型。相比之下,有限差分方法需要扰动输入并重复求值;它们会引入近似误差,在对大量参数分别应用时,计算成本也很高。反向传播直接应用求导规则,不过其数值结果仍会受到浮点舍入的影响。(jmlr.org)
内存是一项重要开销,因为反向计算可能需要前向计算中的中间张量。PyTorch 等系统会记录运算、保存必要的数值,并自动执行相应的反向求导规则。多次调用反向计算时,梯度可能会累积,因此独立的训练步骤之间需要适当地清零梯度。对已保存数值进行原地修改可能使求导失效,而对于不需要导数的计算,关闭梯度记录可以减少开销。(docs.pytorch.org)
历史发展
反向模式微分早于它在神经网络中的广泛应用。塞波·林纳因马1970年的研究通常被视为较早公开发表的描述;保罗·沃博斯1974年的研究也是另一项重要的先驱工作。1986年的论文《通过反向传播误差学习表征》展示了如何利用输出误差的导数训练隐藏层表征,推动了这一方法的普及。因此,反向传播的历史既包括早期的微分研究,也包括后来在机器学习中的应用。(jmlr.org)
局限性
导数的反复相乘可能导致梯度消失问题或梯度爆炸问题。这些效应在循环神经网络中尤为重要,因为其中的依赖关系可能跨越许多时间步。过小的梯度会削弱学习长程依赖所需的信号;过大的梯度则可能使更新不稳定。梯度范数裁剪可以限制过大的梯度,但它本身并不能解决梯度消失问题。(proceedings.mlr.press)
反向传播还要求各个局部运算具有可用的求导规则。在不可微点,软件可能采用指定的处理约定,例如选择一个次梯度;真正的离散运算则可能中断常规的梯度传播。正确的梯度反映的是实际实现的计算,包括其中可能存在的无效算术运算:事后屏蔽一个未定义的结果,并不一定能避免反向计算中出现未定义的梯度。(docs.pytorch.org)