AI 工程基础体系 · 第 58/100 篇。内容覆盖机器学习、深度学习与生成式 AI;模型、数据、评测、权限和成本会作为同一生产系统处理。

自动微分与反向传播:计算图、梯度累积、截断和数值检查

自动微分(automatic differentiation,AD)是按照程序实际执行的运算,利用链式法则计算导数的技术。反向传播(backpropagation)是自动微分在“标量损失由大量参数产生”这一典型机器学习场景下的反向模式(reverse mode)实现。

二者经常被混为一谈:

  • 自动微分回答“如何根据一串可微运算计算导数”;
  • 反向传播回答“当输出主要是一个标量损失时,如何从损失反向计算所有中间变量和参数的梯度”。

PyTorch 的 autograd 就是一个动态自动微分系统。模型执行前向计算时,框架记录计算关系;调用 loss.backward() 时,框架沿着这些关系反向应用链式法则。


一、从导数、梯度和链式法则开始

设标量函数为:

y=f(x)y=f(x)

导数定义为:

dydx=limΔx0f(x+Δx)f(x)Δx\frac{\mathrm d y}{\mathrm d x} = \lim_{\Delta x\to 0} \frac{f(x+\Delta x)-f(x)}{\Delta x}

在机器学习中,变量通常是向量或张量。例如:

y=f(x),xRn,yRm\mathbf y=f(\mathbf x),\qquad \mathbf x\in\mathbb R^n,\quad \mathbf y\in\mathbb R^m

此时一阶导数通常表示为雅可比矩阵:

Jij=yixjJ_{ij} = \frac{\partial y_i}{\partial x_j}

其中第 ii 行描述输出 yiy_i 对所有输入分量的变化,第 jj 列描述某个输入分量对所有输出的影响。

但训练神经网络时,我们通常优化的是一个标量损失:

L=f(θ)L=f(\boldsymbol\theta)

其中 θ\boldsymbol\theta 是所有模型参数。此时对参数的导数构成梯度:

θL=[Lθ1,Lθ2,]\nabla_{\boldsymbol\theta} L = \left[ \frac{\partial L}{\partial \theta_1}, \frac{\partial L}{\partial \theta_2}, \dots \right]

梯度指向函数增长最快的方向,因此最基本的梯度下降更新为:

θt+1=θtηθL\boldsymbol\theta_{t+1} = \boldsymbol\theta_t - \eta\nabla_{\boldsymbol\theta}L

其中 η\eta 是学习率。

1. 链式法则是反向传播的数学基础

如果:

u=g(x),y=f(u)u=g(x),\qquad y=f(u)

那么:

dydx=dydududx\frac{\mathrm d y}{\mathrm d x} = \frac{\mathrm d y}{\mathrm d u} \frac{\mathrm d u}{\mathrm d x}

对多变量函数,链式法则使用雅可比矩阵:

yx=yuux\frac{\partial \mathbf y}{\partial \mathbf x} = \frac{\partial \mathbf y}{\partial \mathbf u} \frac{\partial \mathbf u}{\partial \mathbf x}

深度神经网络只是把许多这样的局部函数串联起来。例如:

h1=ϕ(W1x+b1)\mathbf h_1=\phi(W_1\mathbf x+\mathbf b_1)

h2=ϕ(W2h1+b2)\mathbf h_2=\phi(W_2\mathbf h_1+\mathbf b_2)

L=(h2,y)L=\ell(\mathbf h_2,\mathbf y)

损失对 W1W_1 的导数必须经过:

Lh2W2h1+b2h1W1x+b1W1L \rightarrow \mathbf h_2 \rightarrow W_2\mathbf h_1+\mathbf b_2 \rightarrow \mathbf h_1 \rightarrow W_1\mathbf x+\mathbf b_1 \rightarrow W_1

反向传播就是沿这条依赖路径,从 LL 开始逐步计算局部导数并相乘。


二、自动微分不是数值微分,也不是符号微分

理解自动微分的边界,需要区分三种方法。

1. 符号微分

符号微分直接操作数学表达式。例如:

f(x)=x2+sinxf(x)=x^2+\sin x

可以得到:

f(x)=2x+cosxf'(x)=2x+\cos x

它保留了符号表达式,但复杂程序可能产生极大的表达式,且难以处理动态控制流、循环和张量运算。

2. 数值微分

数值微分使用有限差分近似导数。例如中心差分:

f(x)f(x+ϵ)f(xϵ)2ϵf'(x) \approx \frac{f(x+\epsilon)-f(x-\epsilon)}{2\epsilon}

它不需要知道内部计算过程,但会受到截断误差和浮点舍入误差影响。

3. 自动微分

自动微分把程序拆成一系列基本运算,并对每个基本运算使用已知的局部导数。例如:

  • 加法:(a+b)a=1\frac{\partial(a+b)}{\partial a}=1
  • 乘法:(ab)a=b\frac{\partial(ab)}{\partial a}=b
  • 指数:eaa=ea\frac{\partial e^a}{\partial a}=e^a
  • 正弦:sinaa=cosa\frac{\partial\sin a}{\partial a}=\cos a

它通过链式法则组合这些局部导数。自动微分得到的是“对执行过的程序求导”,不是通过很小的扰动猜测导数,因此通常比数值微分更准确、更适合训练。

不过,自动微分仍然使用浮点数。它不是数学上的无限精度求导,仍可能受到溢出、下溢、舍入误差和非光滑点的影响。


三、计算图:前向执行记录了什么

**计算图(computational graph)**是表示张量之间计算依赖关系的有向图。

如果有:

a=xwa=xw

则可以表示为:

flowchart LR
    x["x"] --> mul["乘法"]
    w["w"] --> mul
    mul --> a["a = x × w"]
    a --> loss["L"]

图中的节点可以代表:

  • 输入张量;
  • 参数张量;
  • 运算结果;
  • 加法、乘法、矩阵乘法、激活函数等运算。

边表示数据依赖关系。若 aa 是由 xxww 计算得到的,则 aa 依赖于 xxww

1. 动态计算图和静态计算图

PyTorch 默认采用动态计算图:Python 代码执行一次,相关的计算关系就被记录一次。

import torch

x = torch.tensor(2.0, requires_grad=True)
w = torch.tensor(3.0, requires_grad=True)

y = x * w
loss = (y - 10.0) ** 2

print(y)       # tensor(6., grad_fn=<MulBackward0>)
print(loss)    # tensor(16., grad_fn=<PowBackward0>)

loss.backward()

print(x.grad)  # tensor(-48.)
print(w.grad)  # tensor(-32.)

这里:

y=xw=2×3=6y=xw=2\times3=6

L=(y10)2=(610)2=16L=(y-10)^2=(6-10)^2=16

根据链式法则:

Ly=2(y10)=8\frac{\partial L}{\partial y} = 2(y-10) = -8

因为:

yx=w=3\frac{\partial y}{\partial x}=w=3

所以:

Lx=Lyyx=8×3=24\frac{\partial L}{\partial x} = \frac{\partial L}{\partial y} \frac{\partial y}{\partial x} = -8\times3=-24

但代码输出是 48-48,这是因为上面的手算遗漏了 y=xwy=xw 的实际数值检查?重新计算可知:

L=(610)2=16L=(6-10)^2=16

Ly=2(610)=8\frac{\partial L}{\partial y}=2(6-10)=-8

yx=w=3\frac{\partial y}{\partial x}=w=3

因此应为:

Lx=24\frac{\partial L}{\partial x}=-24

代码中的 loss = (y - 10.0) ** 2 确实应产生 x.grad=-24w.grad=-16。如果运行结果不是这样,通常说明实际代码与示例不一致。正确的可运行版本如下:

import torch

x = torch.tensor(2.0, requires_grad=True)
w = torch.tensor(3.0, requires_grad=True)

y = x * w
loss = (y - 10.0) ** 2
loss.backward()

print("y =", y.item())             # y = 6.0
print("loss =", loss.item())       # loss = 16.0
print("dL/dx =", x.grad.item())    # dL/dx = -24.0
print("dL/dw =", w.grad.item())    # dL/dw = -16.0

这个例子也说明了一个重要原则:必须把前向值和每一步局部导数一起检查,不能只凭直觉判断反向结果。

2. requires_grad、叶子张量和 grad_fn

当张量设置为 requires_grad=True 时,PyTorch 会追踪由它参与的后续运算。

import torch

x = torch.tensor(2.0, requires_grad=True)
y = x * 3
z = y + 1

print(x.is_leaf)        # True
print(y.is_leaf)        # False
print(y.grad_fn)        # 通常类似 <MulBackward0 ...>
print(z.grad_fn)        # 通常类似 <AddBackward0 ...>

**叶子张量(leaf tensor)**通常是用户直接创建、且需要梯度的参数或输入。默认情况下,反向传播后其梯度保存在 .grad 中。

中间张量通常有 grad_fn,表示它由哪个反向运算节点产生,但默认不把中间张量的梯度保存到 .grad。如果确实需要查看中间梯度,可以显式调用:

x = torch.tensor(2.0, requires_grad=True)
y = x * x
y.retain_grad()

loss = y + 1
loss.backward()

print(y.grad)  # tensor(1.)
print(x.grad)  # tensor(4.)

retain_grad() 只影响中间张量梯度的保存,不改变梯度计算本身。


四、一个完整的反向传播推导

考虑以下标量计算:

a=xwa=xw

b=a+cb=a+c

L=b2L=b^2

取:

x=2,w=3,c=1x=2,\quad w=3,\quad c=1

1. 前向过程

先计算:

a=xw=2×3=6a=xw=2\times3=6

b=a+c=6+1=7b=a+c=6+1=7

L=b2=49L=b^2=49

计算图为:

flowchart LR
    x["x = 2"] --> a["a = xw"]
    w["w = 3"] --> a
    a --> b["b = a + c"]
    c["c = 1"] --> b
    b --> L["L = b²"]

2. 反向过程

从损失开始:

Lb=2b=14\frac{\partial L}{\partial b}=2b=14

因为:

b=a+cb=a+c

所以:

ba=1,bc=1\frac{\partial b}{\partial a}=1,\qquad \frac{\partial b}{\partial c}=1

于是:

La=Lbba=14\frac{\partial L}{\partial a} = \frac{\partial L}{\partial b} \frac{\partial b}{\partial a} = 14

Lc=Lbbc=14\frac{\partial L}{\partial c} = \frac{\partial L}{\partial b} \frac{\partial b}{\partial c} = 14

又因为:

a=xwa=xw

所以:

ax=w=3,aw=x=2\frac{\partial a}{\partial x}=w=3,\qquad \frac{\partial a}{\partial w}=x=2

最终:

Lx=Laax=14×3=42\frac{\partial L}{\partial x} = \frac{\partial L}{\partial a} \frac{\partial a}{\partial x} = 14\times3=42

Lw=Laaw=14×2=28\frac{\partial L}{\partial w} = \frac{\partial L}{\partial a} \frac{\partial a}{\partial w} = 14\times2=28

Lc=14\frac{\partial L}{\partial c}=14

对应代码:

import torch

x = torch.tensor(2.0, requires_grad=True)
w = torch.tensor(3.0, requires_grad=True)
c = torch.tensor(1.0, requires_grad=True)

a = x * w
b = a + c
loss = b ** 2

loss.backward()

print(a.item())       # 6.0
print(b.item())       # 7.0
print(loss.item())    # 49.0
print(x.grad.item())  # 42.0
print(w.grad.item())  # 28.0
print(c.grad.item())  # 14.0

反向传播并不是“从参数重新运行模型并猜测方向”,而是缓存或重新计算前向节点所需的局部信息,再按照依赖关系反向传递伴随量。


五、为什么反向模式适合神经网络

假设模型有 nn 个参数,输出损失是一个标量:

L=f(θ1,,θn)L=f(\theta_1,\ldots,\theta_n)

训练需要的是:

[Lθ1,,Lθn]\left[ \frac{\partial L}{\partial\theta_1}, \ldots, \frac{\partial L}{\partial\theta_n} \right]

反向模式先从一个输出 LL 出发,经过一次反向遍历,就能同时得到所有参数对该标量损失的梯度。因此,对于“参数很多、输出一个标量”的问题,反向模式计算通常更合适。

相对地,前向模式从输入方向传播导数。如果输入维度较少、输出维度较多,前向模式可能更有优势。

可以把两种模式概括为:

  • 前向模式:输入方向 \rightarrow 输出方向;
  • 反向模式:输出方向 \rightarrow 输入方向。

反向传播不是任何自动微分问题的唯一高效方式。它适合标量损失,但如果需要完整的大型雅可比矩阵,直接构造可能非常昂贵,通常应结合 torch.autograd.functionaltorch.func 等工具选择合适的雅可比或向量-雅可比积方法。具体 API 以所使用的 PyTorch 版本文档为准。


六、局部梯度、上游梯度和梯度累积

反向传播中,一个节点通常接收两类信息:

  1. 它对直接输入的局部梯度
  2. 从后续节点传来的上游梯度

二者通过乘法组合。

若:

z=f(x)z=f(x)

并且后续损失为 L(z)L(z),则:

Lx=Lz上游梯度zx局部梯度\frac{\partial L}{\partial x} = \underbrace{\frac{\partial L}{\partial z}}_{\text{上游梯度}} \underbrace{\frac{\partial z}{\partial x}}_{\text{局部梯度}}

1. 分支导致梯度相加

如果一个变量被多个路径使用:

L=L1(x)+L2(x)L=L_1(x)+L_2(x)

则:

Lx=L1x+L2x\frac{\partial L}{\partial x} = \frac{\partial L_1}{\partial x} + \frac{\partial L_2}{\partial x}

例如:

u=x2,v=3x,L=u+vu=x^2,\qquad v=3x,\qquad L=u+v

则:

Lx=2x+3\frac{\partial L}{\partial x}=2x+3

x=2x=2,梯度为 77

import torch

x = torch.tensor(2.0, requires_grad=True)

u = x ** 2
v = 3 * x
loss = u + v

loss.backward()

print(loss.item())  # 10.0
print(x.grad.item())  # 7.0

这就是反向传播中“梯度累积”的第一层含义:同一个变量通过多个路径影响损失时,各路径贡献相加。

2. PyTorch 中 .backward() 默认也会累积梯度

PyTorch 的参数梯度不会在每次 backward() 时自动清零。连续调用会把结果加到已有 .grad 上:

import torch

x = torch.tensor(2.0, requires_grad=True)

loss1 = x ** 2       # d(loss1)/dx = 4
loss1.backward()

print(x.grad.item()) # 4.0

loss2 = 3 * x        # d(loss2)/dx = 3
loss2.backward()

print(x.grad.item()) # 7.0

这与上面的分支求和在数学上是一致的,但代码中两次 backward() 使用了两张独立的计算图。

训练循环通常写成:

for inputs, targets in loader:
    optimizer.zero_grad()
    outputs = model(inputs)
    loss = criterion(outputs, targets)
    loss.backward()
    optimizer.step()

这里 zero_grad() 的作用是清除上一个优化步留下的梯度。若省略它,当前批次的梯度会继续叠加到上一批次,更新方向就不再是当前批次的梯度。

3. 梯度累积训练与忘记清零不是一回事

为了用较小显存模拟大批次,常见做法是将多个小批次的梯度累积后再更新一次参数:

accumulation_steps = 4

optimizer.zero_grad()

for step, (inputs, targets) in enumerate(loader):
    outputs = model(inputs)
    loss = criterion(outputs, targets)

    # 若希望近似“大批次平均损失”,需要除以累积步数
    (loss / accumulation_steps).backward()

    if (step + 1) % accumulation_steps == 0:
        optimizer.step()
        optimizer.zero_grad()

假设每个小批次的损失分别为 L1,,LkL_1,\ldots,L_k,希望得到平均损失:

Lˉ=1ki=1kLi\bar L=\frac1k\sum_{i=1}^{k}L_i

则:

Lˉ=1ki=1kLi\nabla\bar L = \frac1k\sum_{i=1}^{k}\nabla L_i

因此每次反向时使用 Li/kL_i/k,或在最后手动除以 kk,二者在纯梯度累积场景下等价。

但“多个小批次平均梯度”等价于“大批次梯度”需要条件:

  • 各小批次损失的归约方式一致;
  • 损失确实按样本平均,而不是按批次平均后又进行了不一致加权;
  • 模型中没有依赖当前小批次统计量的状态变化,例如某些 BatchNorm 行为;
  • 优化器的更新频率和学习率调度方式相应调整;
  • 随机层、混合精度和分布式同步带来的差异被接受。

4. zero_grad(set_to_none=True) 的含义

常见写法是:

optimizer.zero_grad(set_to_none=True)

它把梯度属性设置为 None,而不是写入全零张量。这样通常可以减少一次写操作,并让框架区分“没有梯度”和“梯度为零”。

二者在普通优化中都能用于清理梯度,但代码若直接访问 .grad,必须处理 None

for parameter in model.parameters():
    if parameter.grad is not None:
        print(parameter.grad.norm())

如果某个参数没有参与当前损失,grad 可能为 None;这与梯度张量恰好全为零是不同状态。


七、计算图的生命周期与重复反向传播

动态计算图通常服务于一次前向和一次反向。反向传播完成后,为节省内存,框架通常会释放反向所需的中间缓存。

因此下面的代码通常会报错:

import torch

x = torch.tensor(2.0, requires_grad=True)
y = x ** 2

y.backward()
y.backward()  # 通常会因计算图中间缓存已释放而报错

如果确实需要对同一张图进行多次反向传播,可以:

x = torch.tensor(2.0, requires_grad=True)
y = x ** 2

y.backward(retain_graph=True)
y.backward()

print(x.grad)  # 两次梯度累积,结果为 16

第一次梯度为 44,第二次又累积 44,所以最终为 88,注意这里若代码中 y=x**2 且 x=2,最终应是 8,不是 16

retain_graph=True 会保留图和中间缓存,适用于确有需要的情况,例如多个目标共享同一前向图。但它会增加内存占用。更常见、更节省内存的方式是重新执行前向计算,建立一张新图。

1. autograd.grad.backward() 的区别

loss.backward() 通常把梯度写入叶子张量的 .grad

torch.autograd.grad 则直接返回指定输入的梯度,不自动把结果累积到这些输入的 .grad

import torch

x = torch.tensor(2.0, requires_grad=True)
y = x ** 3

grad_x, = torch.autograd.grad(y, x)

print(grad_x.item())  # 12.0
print(x.grad)        # None

因为:

dx3dx=3x2=12\frac{\mathrm d x^3}{\mathrm d x}=3x^2=12

当需要计算梯度惩罚、元学习或高阶导数时,autograd.grad 往往更容易控制。


八、截断:计算图截断、梯度截断与截断反向传播

“截断”在深度学习代码中至少有三种不同含义,不能混用。

  1. 计算图截断:切断梯度传播路径,例如 detach()
  2. 梯度裁剪:梯度过大时缩放梯度,例如 clip_grad_norm_
  3. 截断反向传播:在长序列训练中只对固定窗口反向传播,即 TBPTT。

标题中的“截断”主要涉及第一类和第三类,同时需要与梯度裁剪区分。

1. detach():让张量成为不追踪梯度的结果

import torch

x = torch.tensor(2.0, requires_grad=True)
y = x ** 2
z = y.detach()
loss = 3 * z

print(z.requires_grad)  # False

z 的数值与 y 相同,但从 z 继续计算得到的损失不会通过 z 回到 yx

x = torch.tensor(2.0, requires_grad=True)

y = x ** 2
z = y.detach()
loss = 3 * z

print(loss.requires_grad)  # False
# loss.backward() 会报错,因为 loss 没有可反向传播的图

更典型的情况是保留一个数值状态,但不让当前训练图跨越它:

h = torch.zeros(32, requires_grad=True)

for _ in range(10):
    h = cell(h)
    h = h.detach()

每次 detach() 后,新的 h 在数值上延续上一步状态,但梯度不会再回到更早的时间步。

2. torch.no_grad()detach() 的差别

torch.no_grad() 是一个上下文管理器,作用是:在其作用域内,默认不记录新运算的自动微分图。

with torch.no_grad():
    outputs = model(inputs)

这适用于验证和推理,因为不需要训练梯度。

detach() 则是对某个已有张量切断历史。例如:

features = encoder(inputs)
frozen_features = features.detach()
outputs = head(frozen_features)

这里 head 仍可以计算梯度,但梯度不会回到 encoder

冻结参数还涉及 requires_grad 与优化器参数列表。仅仅不把参数加入优化器,不等于前向过程不会构建梯度图;若仍需要对其他部分求梯度,图可能仍会记录相关运算。实际是否关闭参数梯度,应根据“是否要对该模块反向求导”决定。

3. 截断反向传播 TBPTT

循环神经网络或状态空间模型处理长序列时,完整反向传播要沿整个序列保存计算图:

ht=F(ht1,xt)h_t=F(h_{t-1},x_t)

若序列长度为 TT,完整反向需要从 hTh_T 沿:

hThT1h0h_T\rightarrow h_{T-1}\rightarrow\cdots\rightarrow h_0

传播梯度。内存和计算成本会随序列长度增加,且长链路容易出现梯度消失或爆炸。

截断反向传播(Truncated Backpropagation Through Time,TBPTT)把序列拆成长度为 KK 的窗口。窗口边界保留隐藏状态的数值,但切断其历史计算图:

hidden = None

for chunk_inputs, chunk_targets in sequence_chunks:
    if hidden is not None:
        hidden = hidden.detach()

    outputs, hidden = model(chunk_inputs, hidden)
    loss = criterion(outputs, chunk_targets)

    optimizer.zero_grad()
    loss.backward()
    optimizer.step()

其数学含义不是“完整梯度的另一种精确计算”,而是优化一个截断后的目标。若损失为:

L=t=1TtL=\sum_{t=1}^{T}\ell_t

完整梯度包含跨越任意远时间步的依赖:

tθ=stthththt1hs+1hshsθ\frac{\partial \ell_t}{\partial \theta} = \sum_{s\le t} \frac{\partial \ell_t}{\partial h_t} \frac{\partial h_t}{\partial h_{t-1}} \cdots \frac{\partial h_{s+1}}{\partial h_s} \frac{\partial h_s}{\partial\theta}

TBPTT 只保留窗口内的项,窗口外的路径被 detach() 删除。因此:

  • 优点是显著降低图保存的内存压力;
  • 代价是忽略长距离依赖的梯度;
  • 它不是梯度裁剪,也不是把梯度数值变小;
  • 窗口长度 KK 是建模与资源之间的取舍,而非无条件越大越好。

4. 一个常见错误:误用 .data

不应使用:

h = h.data

来替代 detach().data 可能绕过自动微分的版本检查和一致性保护,造成静默错误。应使用:

h = h.detach()

如果需要原地修改,还必须考虑张量是否与计算图中的其他节点共享存储,以及原地操作是否破坏反向所需的值。


九、梯度裁剪不是计算图截断

梯度裁剪发生在反向传播之后、参数更新之前。例如按整体范数裁剪:

loss.backward()

total_norm = torch.nn.utils.clip_grad_norm_(
    model.parameters(),
    max_norm=1.0,
)

optimizer.step()

如果所有梯度拼成向量 gg,其范数为:

g2\|g\|_2

当:

g2>c\|g\|_2>c

则常见裁剪方式将梯度缩放为:

g=gcg2g' = g\cdot \frac{c}{\|g\|_2}

其中 cc 是最大范数。

梯度裁剪:

  • 不删除计算图;
  • 不改变前向值;
  • 不阻止梯度从当前损失回到更早的计算节点;
  • 只改变传给优化器的梯度数值。

detach()

  • 不执行梯度缩放;
  • 直接删除某条反向依赖路径;
  • 使路径上的历史参数收不到来自当前损失的梯度。

可以用下面的因果顺序区分二者:

前向计算
  -> 构建计算图
  -> loss.backward()
  -> 得到梯度
  -> 梯度裁剪(可选)
  -> optimizer.step()

detach() 位于前向图构建阶段,裁剪位于反向完成之后。


十、数值检查:为什么需要有限差分

自动微分的结果通常可信,但复杂模型中仍可能因为以下原因得到错误梯度:

  • 自定义 autograd.Function 的反向公式写错;
  • 张量维度或广播规则理解错误;
  • 在不应使用的地方执行了 detach()
  • 原地操作覆盖了反向需要的值;
  • 混合精度或低精度造成数值问题;
  • 损失归约方式与预期不一致;
  • 只检查了某一部分参数;
  • 非光滑函数处于边界点。

数值梯度检查使用有限差分作为独立参照。对标量函数 f(x)f(x),中心差分为:

gnum(x)=f(x+ϵ)f(xϵ)2ϵg_{\text{num}}(x) = \frac{f(x+\epsilon)-f(x-\epsilon)}{2\epsilon}

自动微分给出:

gauto(x)=fxg_{\text{auto}}(x) = \frac{\partial f}{\partial x}

如果两者接近,说明实现通常没有明显错误。

1. 中心差分的误差来源

中心差分同时存在两类误差:

  1. 截断误差ϵ\epsilon 太大时,局部线性近似不够准确;
  2. 舍入误差ϵ\epsilon 太小时,f(x+ϵ)f(x+\epsilon)f(xϵ)f(x-\epsilon) 的差值可能被浮点精度吞掉。

因此 ϵ\epsilon 不能盲目设为极小值。对于 float32,数值检查通常更容易受到舍入误差影响;检查时常转换到 float64

2. 一个可运行的数值梯度检查

下面检查:

f(x,y)=x2y+sin(x)f(x,y)=x^2y+\sin(x)

解析导数为:

fx=2xy+cos(x)\frac{\partial f}{\partial x}=2xy+\cos(x)

fy=x2\frac{\partial f}{\partial y}=x^2

代码如下:

import math
import torch

torch.set_default_dtype(torch.float64)

def function(x, y):
    return x ** 2 * y + torch.sin(x)

x = torch.tensor(1.2, requires_grad=True)
y = torch.tensor(-0.7, requires_grad=True)

value = function(x, y)
grad_x, grad_y = torch.autograd.grad(value, (x, y))

def finite_difference_x(x_value, y_value, eps=1e-6):
    x_plus = torch.tensor(x_value + eps)
    x_minus = torch.tensor(x_value - eps)
    y_value = torch.tensor(y_value)
    return (
        function(x_plus, y_value) -
        function(x_minus, y_value)
    ) / (2 * eps)

def finite_difference_y(x_value, y_value, eps=1e-6):
    y_plus = torch.tensor(y_value + eps)
    y_minus = torch.tensor(y_value - eps)
    x_value = torch.tensor(x_value)
    return (
        function(x_value, y_plus) -
        function(x_value, y_minus)
    ) / (2 * eps)

numeric_x = finite_difference_x(x.item(), y.item()).item()
numeric_y = finite_difference_y(x.item(), y.item()).item()

print("autograd dx =", grad_x.item())
print("numeric  dx =", numeric_x)
print("autograd dy =", grad_y.item())
print("numeric  dy =", numeric_y)

print(
    "dx close:",
    math.isclose(grad_x.item(), numeric_x, rel_tol=1e-6, abs_tol=1e-8)
)
print(
    "dy close:",
    math.isclose(grad_y.item(), numeric_y, rel_tol=1e-6, abs_tol=1e-8)
)

x=1.2,y=0.7x=1.2,y=-0.7 时:

fx=2(1.2)(0.7)+cos(1.2)0.9793\frac{\partial f}{\partial x} = 2(1.2)(-0.7)+\cos(1.2) \approx -0.9793

fy=1.22=1.44\frac{\partial f}{\partial y} = 1.2^2=1.44

自动微分和数值微分应在合理容差内接近这些结果。

3. 使用 gradcheck

对于张量函数,PyTorch 提供了 torch.autograd.gradcheck,它会在双精度下使用数值差分检查自动微分结果:

import torch
from torch.autograd import gradcheck

def function(inputs):
    x, y = inputs
    return (x ** 2 * y + torch.sin(x)).sum()

inputs = (
    torch.tensor([1.2, -0.4], dtype=torch.float64, requires_grad=True),
    torch.tensor([-0.7, 0.9], dtype=torch.float64, requires_grad=True),
)

ok = gradcheck(function, (inputs,), eps=1e-6, atol=1e-5, rtol=1e-3)
print(ok)  # 通常为 True

使用 gradcheck 时需要注意:

  • 输入通常应为 float64
  • 输入必须设置 requires_grad=True
  • 函数应返回可用于检查的输出;
  • 随机操作会导致同一个输入的两次函数值不一致;
  • 非光滑点、离散操作和条件分支边界可能使检查没有明确意义;
  • 大型神经网络逐个参数做数值检查成本很高,通常只检查小型、确定性的局部模块或自定义算子。

十一、数值检查的反例和边界

1. ReLU 在零点不可导

ReLU 定义为:

ReLU(x)=max(0,x)\operatorname{ReLU}(x)=\max(0,x)

x>0x>0 时导数为 11,当 x<0x<0 时导数为 00。但在 x=0x=0 处,左右导数不同:

limh0ReLU(0+h)ReLU(0)h=0\lim_{h\to0^-}\frac{\operatorname{ReLU}(0+h)-\operatorname{ReLU}(0)}h=0

limh0+ReLU(0+h)ReLU(0)h=1\lim_{h\to0^+}\frac{\operatorname{ReLU}(0+h)-\operatorname{ReLU}(0)}h=1

因此经典导数不存在。框架会选择一个次梯度或约定值,但中心差分可能落在函数两侧,得到约 0.50.5 一类的结果。此时不能简单地把数值差分结果与框架梯度按普通光滑函数比较。

2. 整数张量不支持通常意义上的梯度

x = torch.tensor(2, dtype=torch.int64)

整数变化不是连续变量,不能使用普通实数导数。模型输入可以是整数索引,例如词元 ID,但嵌入层的浮点参数才是被优化的对象。

3. 离散采样通常不可直接反向

argmax、离散采样和硬索引会产生不可微或几乎处处梯度为零的路径。例如:

indices = logits.argmax(dim=-1)

indices 是整数索引,损失通常不能沿它回到 logits。生成式模型中的采样、搜索和离散决策若需要训练,必须采用可微近似、策略梯度或其他专门估计方法,不能假设普通 backward() 会自动穿过离散操作。

4. 低精度会影响梯度检查

float16 或某些 bfloat16 计算中,有限差分的扰动可能小到无法改变表示值,或者函数值计算已经产生较大误差。因此数值检查一般在 float64、CPU、小规模输入上进行,而不是直接在生产训练配置中执行。


十二、梯度消失、梯度爆炸与链式乘积

深层网络和长序列中的梯度问题来自一连串雅可比矩阵或局部导数的乘积。

对递归关系:

ht=F(ht1,xt)h_t=F(h_{t-1},x_t)

有:

hTh0=t=1Ththt1\frac{\partial h_T}{\partial h_0} = \prod_{t=1}^{T} \frac{\partial h_t}{\partial h_{t-1}}

若每一步的有效缩放因子平均小于 11,乘积可能指数级变小,形成梯度消失;若平均大于 11,乘积可能指数级变大,形成梯度爆炸。

这与截断反向传播有关,但不等于截断能解决问题:

  • detach() 减少了反向路径长度,可能缓解内存和部分长链路数值问题;
  • 它同时丢失了窗口之外的真实梯度;
  • 梯度裁剪可以限制爆炸梯度的更新幅度;
  • 它不能恢复已经消失的梯度;
  • 残差连接、归一化、门控结构和合理初始化改变的是梯度传播结构本身。

诊断时应记录而不是猜测:

for name, parameter in model.named_parameters():
    if parameter.grad is not None:
        print(
            name,
            "grad_norm=",
            parameter.grad.detach().norm().item()
        )

如果某层梯度长期为零,需要进一步区分:

  • 参数没有参与当前损失;
  • 参数被 detach()no_grad() 隔断;
  • 激活函数处于零梯度区域;
  • 梯度发生下溢;
  • 损失归约或掩码错误;
  • 参数本身未被优化器管理。

十三、训练、评估和推理中的自动微分状态

训练、评估和推理是不同状态,不能只通过 model.train()model.eval() 其中一个开关概括。

1. model.train()model.eval()

它们主要控制 Dropout、BatchNorm 等模块的行为:

model.train()  # 训练行为
model.eval()   # 评估行为

它们不等价于打开或关闭自动微分。

2. 验证阶段通常使用 torch.no_grad()

model.eval()

with torch.no_grad():
    for inputs, targets in validation_loader:
        outputs = model(inputs)
        loss = criterion(outputs, targets)

这样既使用评估模式,又不保存梯度图,减少内存和计算开销。

如果需要在评估阶段计算输入梯度,例如对抗样本、显著性分析或某些解释任务,则不能使用 no_grad() 包住相关计算:

model.eval()

inputs = inputs.detach().requires_grad_(True)
outputs = model(inputs)
loss = criterion(outputs, targets)

input_grad, = torch.autograd.grad(loss, inputs)

这段代码只对输入求梯度,不把结果累积到模型参数的 .grad

3. torch.inference_mode()

torch.inference_mode() 是面向纯推理场景的更强优化上下文。它适合不需要构建自动微分图、也不需要在后续重新把结果用于梯度计算的推理路径。

如果某个中间结果之后仍需要参与梯度计算,应使用普通的梯度上下文,而不是不加区分地使用推理模式。具体限制与行为以当前 PyTorch 文档为准。


十四、高阶导数需要保留反向图

第一次反向传播通常生成一阶梯度。如果还要对梯度再次求导,就需要在计算一阶梯度时保留其计算过程:

import torch

x = torch.tensor(3.0, requires_grad=True)

y = x ** 3
first_grad, = torch.autograd.grad(
    y,
    x,
    create_graph=True,
)

second_grad, = torch.autograd.grad(first_grad, x)

print(first_grad.item())   # 27.0
print(second_grad.item())  # 18.0

因为:

y=x3y=x^3

dydx=3x2\frac{\mathrm dy}{\mathrm dx}=3x^2

d2ydx2=6x\frac{\mathrm d^2y}{\mathrm dx^2}=6x

x=3x=3 时,一阶导数为 2727,二阶导数为 1818

create_graph=True 会把一阶梯度计算也纳入新的计算图,因而增加内存和计算成本。它不是普通训练循环的默认选项,只在梯度惩罚、元学习、二阶优化等场景需要。


十五、自定义自动微分函数的正确性

当实现自定义算子时,需要同时明确:

  1. 前向函数计算什么;
  2. 反向函数接收什么上游梯度;
  3. 反向函数返回每个输入的梯度;
  4. 对不需要梯度的输入返回 None
  5. 保存哪些前向值供反向使用。

例如:

y=x2y=x^2

其反向关系为:

Lx=Ly2x\frac{\partial L}{\partial x} = \frac{\partial L}{\partial y}\cdot 2x

在自定义反向中,不能只返回局部导数 2x2x,还必须乘以上游梯度 L/y\partial L/\partial y。这是自定义反向实现中非常常见的错误。

自定义函数还必须考虑:

  • 广播后的梯度形状是否需要求和还原;
  • 输入是否是非连续张量;
  • 原地修改是否破坏保存的前向值;
  • float32float64、复杂数等类型;
  • 二阶导数是否仍然正确;
  • 随机操作是否能稳定进行梯度检查。

gradcheck 适合验证这类局部实现,但通过 gradcheck 只说明被覆盖的输入和路径在给定条件下匹配数值导数,并不证明整个模型训练目标正确。


十六、常见失败表现与诊断路径

1. element 0 of tensors does not require grad

常见原因是损失没有连接到任何需要梯度的张量,例如:

with torch.no_grad():
    loss = criterion(model(x), target)
loss.backward()

或中间结果被完全 detach()

loss = criterion(model(x).detach(), target)

诊断:

print(loss.requires_grad)
print(loss.grad_fn)

正常的训练损失通常应有 requires_grad=True,并且存在 grad_fn。如果二者都没有,说明图在损失之前已经断开,或模型参数不需要梯度。

2. 梯度为 None

可能原因包括:

  • 参数没有参与当前损失;
  • 参数被排除在计算路径之外;
  • 使用了 set_to_none=True 且尚未产生梯度;
  • 参数被冻结;
  • 前向分支没有执行。

不要把 None 自动当成零。应先确认该参数是否理论上应参与当前损失。

3. 梯度数值不断叠加

如果每个优化步前没有清理梯度:

loss.backward()
optimizer.step()

那么 .grad 会跨步累积。若这是有意的梯度累积训练,应明确累积窗口、损失缩放和更新时机;否则应在每次更新前调用 zero_grad()

4. 第二次 backward() 报图已释放

通常是同一张计算图被重复反向。解决方式不是无条件添加 retain_graph=True,而是先判断需求:

  • 需要独立训练步:重新前向;
  • 需要同一图多次反向:使用 retain_graph=True,接受额外内存;
  • 只是要多个输入的梯度:考虑一次 autograd.grad 指定多个输入。

5. 反向传播时出现原地操作错误

某些前向值被反向过程保存。如果在图仍然需要它时原地修改,框架可能报告变量版本不匹配,或者在不安全场景下产生错误结果。

优先使用非原地表达式:

x = x + residual

而不是在不确定依赖关系时使用:

x += residual

原地操作是否安全取决于具体算子和张量别名关系,不能仅凭“数值看起来一样”判断。

6. 数值梯度检查差异很大

应按以下顺序缩小问题:

  1. 使用无随机性的最小输入;
  2. 切换到 float64
  3. 检查输出是否有限:torch.isfinite
  4. 尝试多个 ϵ\epsilon,例如 10410^{-4}10510^{-5}10610^{-6}
  5. 排除 ReLU 零点、argmax 和离散操作;
  6. 比较单个参数或单个输入分量;
  7. 检查广播、求和归约和掩码;
  8. 再将验证扩展到完整模块。

若函数本身含随机性,数值差分的两次前向必须使用一致的随机状态,否则比较的是两个不同函数值。


十七、一个更完整的训练循环示例

下面的示例展示了训练中计算图、梯度清零、反向传播、梯度范数检查和参数更新的顺序:

import torch
from torch import nn

torch.manual_seed(0)

model = nn.Sequential(
    nn.Linear(4, 8),
    nn.ReLU(),
    nn.Linear(8, 1),
)

criterion = nn.MSELoss()
optimizer = torch.optim.AdamW(model.parameters(), lr=1e-3)

model.train()

for step in range(3):
    inputs = torch.randn(16, 4)
    targets = torch.randn(16, 1)

    # 清除上一个优化步留下的梯度
    optimizer.zero_grad(set_to_none=True)

    # 前向:动态创建本轮计算图
    predictions = model(inputs)
    loss = criterion(predictions, targets)

    # 反向:沿图计算并累积参数梯度
    loss.backward()

    # 在更新前检查梯度是否有限
    all_finite = True
    for parameter in model.parameters():
        if parameter.grad is not None:
            all_finite = all_finite and bool(
                torch.isfinite(parameter.grad).all()
            )

    if not all_finite:
        raise FloatingPointError("检测到非有限梯度")

    # 可选:限制异常大的梯度
    grad_norm = nn.utils.clip_grad_norm_(
        model.parameters(),
        max_norm=1.0,
    )

    # 根据梯度更新参数
    optimizer.step()

    print(
        f"step={step}, "
        f"loss={loss.item():.6f}, "
        f"grad_norm={float(grad_norm):.6f}"
    )

每一步的因果关系是:

  1. zero_grad 防止上一步梯度混入当前优化步;
  2. 前向计算建立本轮动态图;
  3. backward 计算当前损失对参数的梯度;
  4. isfinite 检查是否出现 NaN 或无穷大;
  5. 梯度裁剪只改变传给优化器的梯度;
  6. optimizer.step() 修改参数;
  7. 下一轮重新建立一张新的计算图。

如果在 optimizer.step() 之后还要使用上一轮图中的中间结果进行反向,必须注意参数已经被原地更新,相关图可能不再适用。生产训练循环通常在一次更新完成后丢弃该图和中间张量。


十八、在生成式 AI 系统中的具体体现

在 Transformer 训练中,模型通常计算:

L=t=1Tlogpθ(xtx<t)L = -\sum_{t=1}^{T} \log p_\theta(x_t\mid x_{<t})

反向传播需要从所有位置的损失回到共享的嵌入矩阵、注意力层、前馈层和输出投影层。一个参数可能被:

  • 多个序列位置共享;
  • 多个样本共享;
  • 输入嵌入和输出投影权重共享;
  • 多条残差路径共同使用。

因此梯度累积是结构性的:来自不同 token、样本和分支的梯度贡献必须求和。

在长上下文训练中,通常直接对完整 token 序列反向传播会保存大量激活。工程上会使用:

  • micro-batch 梯度累积:减少单卡峰值显存;
  • activation checkpointing:重新计算部分前向以减少保存的激活;
  • 序列分块或状态截断:限制反向路径;
  • 混合精度:降低激活和梯度存储成本。

这些机制作用不同:

  • 梯度累积改变参数更新的批次组织;
  • checkpointing 通常保留数学上的反向路径,但用额外前向计算换内存;
  • detach() 或 TBPTT 删除跨边界梯度;
  • 混合精度改变数值表示,需要额外处理溢出和缩放。

在推理阶段,通常不需要反向图;在微调、偏好优化、奖励模型训练或输入梯度分析中,则必须明确哪些路径需要保留,哪些路径有意截断。权限、数据来源、模型版本和成本控制不会改变链式法则,但会改变哪些计算实际被执行、保存和重复,从而影响系统资源与审计结果。


十九、需要牢牢记住的边界

自动微分保证的是:在框架支持的运算和给定执行路径上,按照局部导数与链式法则计算梯度。它不保证:

  • 模型目标函数设计正确;
  • 张量维度和广播符合业务意图;
  • 离散决策可以直接求导;
  • 非光滑点具有唯一梯度;
  • 低精度训练不会溢出;
  • detach() 没有误切断必要路径;
  • 梯度累积后的更新等价于某个理想大批次;
  • 自定义反向函数一定正确;
  • 训练、验证和推理使用了正确的状态。

因此,可靠的梯度流程应同时检查三件事:

  1. 数学路径:损失通过哪些运算依赖哪些参数;
  2. 框架状态:是否启用了梯度追踪,是否意外调用了 detach()no_grad() 或原地操作;
  3. 数值证据:梯度是否有限,规模是否合理,小型可控函数是否通过有限差分或 gradcheck

当这三层都能解释时,loss.backward() 就不再是一个黑盒调用,而是一个可以沿计算图逐节点追踪、用链式法则推导、用数值方法验证,并在资源约束下有意识截断或累积的工程机制。


系列导航与关联阅读

官方资料

本文依据研究论文、标准组织与主流框架官方文档重新梳理;正文、示例与工程清单由 WR BLOG 编写。