形式陈述
自动微分的输入是一段求值程序、求导点以及一个种子向量;输出是该程序所表示映射在此点的导数作用于种子的结果。设程序在实数算术下计算 F : R n → R m ,输入节点为 z j = x j (1 ≤ j ≤ n ),随后按顺序执行 T 个原语:
z i = ϕ i ( z p ( i , 1 ) , … , z p ( i , k i ) ) , p ( i , r ) < i , n < i ≤ n + T . 输出是指定节点组成的列向量 F = ( z o 1 , … , z o m ) 。每个 p ( i , r ) 指一个操作数槽 ;两个槽可以指向同一节点,例如 u ⋅ u 。这些依赖组成有向无环计算图,拓扑顺序就是合法执行顺序。假设各原语在实际到达的操作数处可微,且相应局部复合在输入的某个开邻域内有定义。含分支的程序还须在该邻域内保持同一执行迹,才能直接使用本页定理。
记局部偏导数为
c i , r = ∂ ϕ i ∂ z p ( i , r ) ; 这里偏导是对第 r 个形式参数求导,再代入原值,不能先把相同节点的槽合并。由链式法则 公理库 链式法则 Chain rule 复合映射的导数等于各层导数按计算顺序组成的线性映射复合。 ,下面两种传播分别计算Jacobian 矩阵 公理库 Jacobian 矩阵 Jacobian matrix 多元映射各偏导数组成并表示其导数的矩阵。 的向量积和转置向量积,且都不要求显式形成完整矩阵。
前向模式:输入一个方向
给定 v ∈ R n ,初始化 z ˙ j = v j ,然后按原求值顺序,同时计算原值与切向量:
z ˙ i = ∑ r = 1 k i c i , r z ˙ p ( i , r ) . 返回 ( z ˙ o 1 , … , z ˙ o m ) = J F ( x ) v ,称为 Jacobian-vector product(JVP)。输入方向只有一个,输出可以有多个分量。
反向模式:输入一个输出权重
给定 w ∈ R m ,先运行原程序,保存反向所需的原值与依赖记录,通常称为 tape。把全部伴随量 z ¯ i 置零,对每个输出槽执行 z ¯ o ℓ + = w ℓ ;输出节点若重复出现也要累加。随后按 i = n + T , … , n + 1 的逆序,对每个操作数槽执行
z ¯ p ( i , r ) + = z ¯ i c i , r . 节点 i 的所有贡献分发完后即可把它标记为已处理。最后返回输入节点的伴随量 ( z ¯ 1 , … , z ¯ n ) = J F ( x ) T w 。通常所说的 vector-Jacobian product(VJP)写为行向量 w T J F ( x ) ;本页统一返回它的转置。
正确性:两种保持不变的含义
前向不变量是:每个已经处理的节点都满足 z ˙ i = D z i ( x ) [ v ] 。输入节点显然满足此式。若所有父节点满足,链式法则便把 D ϕ i 作用于这些切向量,恰好得到上述求和更新。沿拓扑顺序归纳,输出就是 J F ( x ) v 。
反向不变量是一个微分线性形式。固定 w ,令 λ = w T F ;初始化后有
d λ = ∑ j z ¯ j d z j . 把尚未消去的节点暂时视为形式变量。逆序处理 z i 时,所有使用它的后继节点都已消去,因此其当前系数已收齐。将
d z i = ∑ r c i , r d z p ( i , r ) 代入线性形式,删去 z ¯ i d z i ,并向每个父节点原有系数加入 z ¯ i c i , r ,线性形式保持不变。最终只剩输入微分,故
d λ = ∑ j = 1 n z ¯ j d x j = w T J F ( x ) d x . 比较系数便得到 x ¯ = J F ( x ) T w 。实现可以保留已处理节点的伴随值以供查看,但它已不在尚待消去的线性形式中。这个证明也解释了为什么更新必须是加法:同一变量对多个后继的影响,以及同一原语内重复槽的影响,都属于同一个微分的系数。
直觉
前向模式问:“沿指定输入方向轻推一下,每个中间量怎样变化?”一个数旁边携带一个切向量分量,程序每执行一个原语,就顺手执行它的一阶变化规则。它沿数据流传递扰动,适合只关心少数输入方向的情形。
反向模式问:“最终加权输出对这个中间量有多敏感?”它先知道全部中间原值,再从输出往回分配敏感度。共享节点像多条依赖路径的汇合处,必须等所有消费者的贡献都到齐,才能继续向输入分配。所谓反向传播,就是用局部链式法则系统地完成这次汇总。
原值、切向量和伴随量含义不同。u ˙ 描述输入扰动引起的 u 的变化;u ¯ 描述目标对 u 的敏感度。二者可能偶然相等,但不能互换。前向没有选定标量目标,反向则由输出权重 w 明确指定目标 w T F 。
例子与边界
共享节点与重复槽的完整执行
取如下四步程序,在 ( x , y ) = ( 2 , 3 ) 处求导:
u = x y , a = u ⋅ u , b = u + x , f = a + b . 原值依次为 u = 6 , a = 36 , b = 8 , f = 44 。选输入方向 v = ( 1 , − 1 ) ,前向计算如下。
节点
原值
切向量计算
切向量
x
2
输入种子
1
y
3
输入种子
− 1
u = x y
6
y x ˙ + x y ˙ = 3 − 2
1
a = u ⋅ u
36
u u ˙ + u u ˙ = 6 + 6
12
b = u + x
8
u ˙ + x ˙ = 1 + 1
2
f = a + b
44
a ˙ + b ˙ = 12 + 2
14
反向令 f ¯ = 1 ,其余伴随为零,按 f , b , a , u 处理。f 向 a , b 各送去 1 ;b 向 u , x 各送去 1 ;a 的两个乘法槽分别向 u 送去 6 。因此在处理 u 之前,u ¯ = 1 + 6 + 6 = 13 ,且 x ¯ = 1 。最后 u = x y 向 x 送去 13 y = 39 ,向 y 送去 13 x = 26 ,得到
∇ f ( 2 , 3 ) = ( 40 , 26 ) T , ∇ f ( 2 , 3 ) T ( 1 , − 1 ) = 14. 图片加载失败 图中共享节点 u 的伴随量由三次贡献相加得到。上半图节点标出原值与切向量;下半图标出最终伴随量,箭头沿反向传播方向,边上数字是本次送出的贡献。u ⋅ u 的两条弧分别代表两个操作数槽,不能只保留其中一条。
独立展开 f = x 2 y 2 + x y + x ,可得 ∂ x f = 2 x y 2 + y + 1 = 40 ,∂ y f = 2 x 2 y + x = 26 。这既核对了结果,也定位了覆盖累加的错误:若把 u ¯ 每次赋成新贡献,来自其他路径的信息就会丢失。
若将输出改为 F = ( a , b ) ,则
J F ( 2 , 3 ) = ( 36 24 4 2 ) , J F ( 2 , 3 ) ( 1 , − 1 ) T = ( 12 , 2 ) T , 而任意输出种子给出
J F ( 2 , 3 ) T w = ( 36 w 1 + 4 w 2 , 24 w 1 + 2 w 2 ) T . w = ( 1 , 1 ) 恢复 f = a + b 的梯度;w = ( 1 , 0 ) 只得到第一行的转置。一次反向返回一个加权行组合,并没有同时返回任意完整 Jacobian。
执行迹和浮点数的边界
分支被选中并不意味着该分支的导数就是整体函数的导数。程序“若 x = 0 则返回 0 ,否则返回 x ”在实数上表示的函数就是 f ( x ) = x ,零点的真实导数为 1 ;但在零点执行的常数返回迹可给出导数 0 。问题在于没有一个零点邻域始终执行这条迹。若分支在邻域内稳定,便可对该迹应用前面的定理;循环同样需要把有限、局部稳定的执行迹展开后分析。
| x | 或 ReLU 在零点没有经典导数。框架在此返回约定值可以服务某种优化算法,却不会让“各原语可微”的假设成立。错误的自定义求导规则也会破坏局部链式法则,传播过程本身无法纠正它。
自动微分采用实数原语的导数规则,再用浮点运算执行这些规则;它不是离散浮点输入输出映射的数学导数。与差分求导的截断—舍入权衡 公理库 数值微分中的截断—舍入权衡 Numerical differentiation and roundoff · Step-size selection for finite differences 以截断项和函数值舍入项的竞争解释差分误差的 U 形曲线,并据此选择步长和失败诊断。 相比,它没有差分步长带来的截断项,但局部乘加仍受浮点误差模型 公理库 浮点算术标准误差模型 Standard floating-point arithmetic model 以每次基本运算的小相对扰动和 gamma 记号组织多步浮点误差分析。 约束。舍入、溢出、巨大中间导数和相消仍可使结果不可靠,“没有步长”并不等于“没有误差”。
推论与应用
工作量、存储与模式选择
假设每个原语的元数有统一常数上界,原值和全部局部偏导都能以有界成本计算。每个槽在一次传播中只访问常数次,所以一个前向方向和一个反向输出种子的传播工作量都是 O ( T ) ;反向所需的原程序求值也在同一量级。若把种子初始化和读写结果纳入成本,通用实现还需计入 O ( n + m ) 。矩阵乘法等大原语必须按实际算术量计费,不能因为写成一条调用就视为常数成本。
前向只需保存仍会被使用的原值及切向量;若最大同时存活量为 L ,额外工作存储为 O ( L ) ,通常不是 O ( 1 ) 。朴素反向保存 tape 和全部伴随,存储为 O ( T + n + m ) 。检查点方法可丢弃部分原值、在反向需要时重算,从而用额外工作换取更少存储;这项权衡不改变反向累加的正确性条件。
对完整稠密 m × n Jacobian,前向用 n 个坐标种子求各列,传播约需 O ( n T ) ;反向用 m 个坐标种子求各行,传播约需 O ( m T ) ,各次可复用原值 tape,但须重新初始化伴随。两者还须支付写出 n m 个数的成本。因而输入方向很少时优先考虑前向,输出目标很少时优先考虑反向;大量参数对应一个标量损失的情形,正是反向模式的典型用途。
在上述四步例子中,原值求值是两次乘法、两次加法。按逐槽规则执行,不合并平方的两个槽、不省略零项,但省略与加法原语的单位局部导数相乘的运算,前向另需四次乘法、四次加法;反向另需四次乘法、八次累加。这些是该图的具体计数。不同原语、融合运算与编译优化会改变常数,不能据此声称所有程序都有相同倍数开销。
从结果到可检查的求导流程
实际选择可从所需对象开始:只需 J v 时运行前向;只需 J T w 时运行反向;需要完整矩阵时才组织多种子求值。标量损失的反向传播输出全部输入梯度,灵敏度分析则常只需要少数指定方向,不必为未使用的偏导支付矩阵存储成本。
检查实现时,可比较 w T ( J v ) 与 ( J T w ) T v ,这会暴露部分种子、转置或累加错误,但两条路径若共享同一个错误局部导数,仍可能相等。再结合可解析的小例子与多步长差分检查,才能分别核对传播规则和局部求导规则。本例的 14 = 40 − 26 、显式多项式偏导以及逐节点表,共同提供了可以复算的检查终点。
参考资料
Thomas Reps, Automatic Differentiation and Backpropagation , CS701 lecture notes, 2015-12-01,§§2–5 的程序、计算图与路径贡献,§6 的反向传播。
Mike Giles, Numerical Methods II , Lecture 16 ,slides 10–15 的状态扩张、前向与反向传播及存储需求。
Walter Baur and Volker Strassen, “The Complexity of Partial Derivatives” , Theoretical Computer Science 22, 1983, pp. 317–330,§1 的代数复杂度模型与 §2, Theorem 1。文中的常数界采用特定的有理函数运算计费模型,不是任意浮点程序的统一倍数界;本页的 O ( T ) 结论由有界原语逐槽计数直接得到。