上一篇 Transformer 入门 把训练循环概括成“前向计算—计算损失—反向传播—更新参数”。这四步看起来像一句固定口令,却连接了三个容易混在一起的问题:梯度是什么,反向传播怎样高效得到梯度,优化器又怎样使用梯度。

本文从只有两个参数的直线开始,先手算一次前向与反向,再把同一过程推广成计算图和自动微分。最后从零实现一个小型反向模式自动微分引擎,用 SGD 与 Adam 学回 y = 2x - 1。目标不是复刻完整深度学习框架,而是建立一条可以逐项检查的训练链。

1. 先分清三件事

给定参数 θ 和损失 L(θ),训练想让损失变小。围绕这件事有三个不同层次:

名称 回答的问题 产物
微分 一个微小参数变化会怎样影响损失 导数或梯度
反向传播 复合计算中怎样复用链式法则 从输出到各参数的梯度
优化器 得到梯度后怎样改变参数 下一步参数值与优化器状态

反向传播不是优化器,Adam 也不负责推导导数。自动微分系统先按程序实际执行的运算构建依赖关系,再在反向阶段计算梯度;优化器读取这些梯度,决定更新的方向与尺度。

2. 梯度是一张局部地图

一元函数的导数描述当前位置附近的斜率。参数是向量时,梯度把损失对每个参数的偏导排成同形状的向量:

1
∇L(θ) = [∂L/∂θ₁, ∂L/∂θ₂, ..., ∂L/∂θₙ]

若学习率为 η,最基本的梯度下降更新为:

1
θ ← θ - η∇L(θ)

负梯度是在当前点的一阶近似下下降最快的方向,但这只是一张局部地图。学习率过大可能一步跨过低谷,过小则进展缓慢;在神经网络的非凸损失面上,它也不承诺抵达全局最优点。

梯度现象 它说明什么 它没有说明什么
梯度为正 增大该参数会在局部增大损失 参数必须永远减小
梯度绝对值大 当前局部敏感度较高 一定应走同样大的步长
梯度接近 0 局部一阶变化小 已找到理想解
不同参数尺度差异大 损失面对各方向的尺度不同 自适应优化器一定能完全解决

3. 手算一条只有两个参数的训练链

用直线预测:

1
2
ŷ = wx + b
L = 1/2 · (ŷ - y)²

x = 2,w = 3,b = -1,y = 4。前向计算得到 ŷ = 5,误差 e = ŷ - y = 1,损失 L = 0.5

反向时从损失开始,沿计算的相反方向应用链式法则:

1
2
3
4
5
∂L/∂ŷ = ŷ - y = 1
∂ŷ/∂w = x = 2
∂ŷ/∂b = 1
∂L/∂w = ∂L/∂ŷ · ∂ŷ/∂w = 2
∂L/∂b = ∂L/∂ŷ · ∂ŷ/∂b = 1

学习率取 0.1,一步更新后 w = 2.8,b = -1.1。同一个样本上的新预测为 4.5,损失降到 0.125

前向值 局部导数 收到的上游梯度 向下游传出的梯度
损失 L = e²/2 0.5 ∂L/∂e = 1 1 1
误差 e = ŷ-y 1 ∂e/∂ŷ = 1 1 1
预测 ŷ = wx+bw 5 ∂ŷ/∂w = 2 1 2
预测 ŷ = wx+bb 5 ∂ŷ/∂b = 1 1 1

这张表里的“上游梯度 × 局部导数”就是反向传播反复做的事情。

4. 为什么要把程序看成计算图

计算图中的节点保存中间结果,边表示依赖关系。前向阶段按拓扑顺序得到各节点的值;反向阶段从标量损失的梯度 1 出发,逆拓扑遍历节点。

一个中间量可能流向多条路径。例如 a 同时参与 a × ba + c,总梯度必须把两条路径的贡献相加:

1
∂L/∂a = 来自第一条路径的贡献 + 来自第二条路径的贡献

这也是框架中梯度默认累积而非覆盖的原因。它既支持分支计算图,也支持把多个微批次的梯度累加后再更新;但若每个训练步本应独立,忘记清零就会把历史梯度意外带入下一步。

图中的结构 反向规则 常见错误
串联 a → b → L 沿路径连乘局部导数 漏掉中间一层链式因子
分支 a → b,c → L 各路径贡献相加 后到的分支覆盖先到的梯度
参数被多次复用 所有使用位置的贡献相加 误以为共享参数只算一次局部导数
不参与损失的节点 对损失的梯度为 0 或不存在 把未连接参数当成数值错误

5. 反向模式为什么适合神经网络

自动微分不等于符号求导,也不等于有限差分。

方法 怎样得到导数 优点 主要限制
符号求导 操作代数表达式 可得到显式公式 表达式可能迅速膨胀,难直接跟随普通程序控制流
数值差分 轻微扰动输入并重复前向 实现简单,适合核验 每个参数都要额外计算,受步长与浮点误差影响
自动微分 对基本算子应用已知导数并组合 精确到浮点运算,可复用中间量 需要记录图或生成导数程序,并保存必要状态

神经网络通常有大量参数,最终损失却是一个标量。反向模式从一个输出向许多输入传播,一次反向过程就能得到所有参数对损失的梯度,因此比逐参数扰动高效得多。

更严格地说,每个反向节点接收的是“某个标量目标对当前节点输出的梯度”,再计算向量—雅可比积(VJP),而不是到处显式构造完整雅可比矩阵。理解这一点有助于解释:非标量输出调用反向时,为什么需要提供一个上游梯度。

6. 动态图、叶子节点与保存的中间量

PyTorch 的 autograd 是反向模式自动微分系统。启用梯度记录的张量参与运算时,系统在前向过程中记录实际执行的操作,并在每轮重新建立图,因此普通 Python 分支和循环可以改变本轮图结构。

框架不会为反向“记住所有东西”,而是由各算子保存计算其导数所需的值。例如平方的反向需要前向输入;某些算子可能保存输出或其他上下文。这些保存值带来训练显存开销,也是激活检查点用额外重算换显存的出发点。

PyTorch 概念 含义 排查时看什么
requires_grad 是否追踪与该张量有关的梯度计算 冻结参数是否真的关闭记录
叶子张量 用户直接创建、没有生成它的 grad_fn 的张量 参数梯度通常累积在叶子的 .grad
grad_fn 产生非叶子张量的反向节点入口 结果是否意外脱离计算图
backward() 从目标触发反向传播并累积叶子梯度 是否对非标量目标提供上游梯度
no_grad 当前代码块不记录反向图 参数更新与纯评估是否需要记录

model.eval() 与关闭梯度是两件事:前者切换 Dropout、BatchNorm 等模块的训练/评估行为,后者决定是否记录自动微分图。验证时通常同时设置评估模式与 no-grad 或 inference mode,但不能把它们当成同一个开关。

7. SGD、动量与 Adam 各自改了什么

7.1 随机梯度下降

全量梯度下降用整个数据集计算一次梯度;随机梯度下降通常泛指用单样本或小批次得到带噪声的梯度估计。最基本的 SGD 直接更新:

1
θₜ = θₜ₋₁ - ηgₜ

小批次越大,梯度估计通常更稳定,但单步计算和内存需求也更高。批大小改变时,有效学习率、归一化层行为和训练步数都可能需要重新考虑。

7.2 动量

动量累计过去梯度的指数衰减方向,减少来回振荡并在方向持续一致时加速。不同资料对速度变量的符号和缩放约定可能不同,比较公式时应先核对实现。

1
2
vₜ = μvₜ₋₁ + gₜ
θₜ = θₜ₋₁ - ηvₜ

7.3 Adam

Adam 同时维护梯度的一阶矩估计 m 与平方梯度的二阶原始矩估计 v,并修正它们从 0 初始化造成的早期偏差:

1
2
3
4
5
mₜ = β₁mₜ₋₁ + (1-β₁)gₜ
vₜ = β₂vₜ₋₁ + (1-β₂)gₜ²
m̂ₜ = mₜ / (1-β₁ᵗ)
v̂ₜ = vₜ / (1-β₂ᵗ)
θₜ = θₜ₋₁ - ηm̂ₜ / (√v̂ₜ + ε)

Adam 为不同参数坐标形成不同的有效步长,常能较快进入可用区域,但“自适应”不代表无需选择学习率,也不保证泛化一定优于 SGD。

优化器 每个参数额外状态 更新直觉 需要留意
SGD 跟随当前批次梯度 对学习率与尺度较敏感
SGD + momentum 一份速度 平滑历史方向 动量约定、学习率调度
Adam 一阶矩与二阶矩 平滑方向并按历史平方梯度缩放 状态内存、βε 与学习率
AdamW Adam 状态 将权重衰减与损失梯度更新解耦 哪些参数不应衰减,如偏置和部分归一化参数

8. 权重衰减不总等于把 L2 项塞进损失

对普通 SGD,在损失中加入 λ‖θ‖²/2 会给梯度增加 λθ,与每步按比例缩小参数有紧密对应关系。对 Adam 这类会按坐标缩放梯度的优化器,把 L2 梯度混入自适应缩放,与直接衰减参数并不等价。

AdamW 将权重衰减从梯度更新中解耦。工程上还要决定哪些参数组使用衰减,而不是看到“AdamW”就认为正则化已经自动正确。优化器状态、参数分组和学习率调度共同决定实际更新。

9. 从零实现一个最小自动微分引擎

下面只用 Python 标准库实现标量节点、加法、乘法、幂、逆拓扑反向传播、SGD 和 Adam。它用 5 个点拟合 y = 2x - 1,适合理解机制,不适合替代张量库。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
136
import json
from math import sqrt


class Value:
def __init__(self, data, parents=(), backward=None):
self.data = float(data)
self.grad = 0.0
self.parents = tuple(parents)
self._backward = backward or (lambda: None)

@staticmethod
def wrap(value):
return value if isinstance(value, Value) else Value(value)

def __add__(self, other):
other = self.wrap(other)
out = Value(self.data + other.data, (self, other))

def backward():
self.grad += out.grad
other.grad += out.grad

out._backward = backward
return out

__radd__ = __add__

def __mul__(self, other):
other = self.wrap(other)
out = Value(self.data * other.data, (self, other))

def backward():
self.grad += other.data * out.grad
other.grad += self.data * out.grad

out._backward = backward
return out

__rmul__ = __mul__

def __pow__(self, exponent):
if not isinstance(exponent, (int, float)):
raise TypeError("exponent must be a number")
out = Value(self.data ** exponent, (self,))

def backward():
self.grad += exponent * (self.data ** (exponent - 1)) * out.grad

out._backward = backward
return out

def __neg__(self):
return self * -1

def __sub__(self, other):
return self + (-self.wrap(other))

def __rsub__(self, other):
return self.wrap(other) + (-self)

def __truediv__(self, other):
return self * (self.wrap(other) ** -1)

def backward(self):
order = []
visited = set()

def build(node):
if node in visited:
return
visited.add(node)
for parent in node.parents:
build(parent)
order.append(node)

build(self)
self.grad = 1.0
for node in reversed(order):
node._backward()


def zero_grad(parameters):
for parameter in parameters:
parameter.grad = 0.0


def sgd_step(parameters, learning_rate):
for parameter in parameters:
parameter.data -= learning_rate * parameter.grad


def adam_step(parameters, state, step, learning_rate=0.1,
beta1=0.9, beta2=0.999, epsilon=1e-8):
for parameter, slot in zip(parameters, state):
slot["m"] = beta1 * slot["m"] + (1 - beta1) * parameter.grad
slot["v"] = beta2 * slot["v"] + (1 - beta2) * parameter.grad ** 2
corrected_m = slot["m"] / (1 - beta1 ** step)
corrected_v = slot["v"] / (1 - beta2 ** step)
parameter.data -= learning_rate * corrected_m / (sqrt(corrected_v) + epsilon)


def train(optimizer, steps):
samples = [(-2, -5), (-1, -3), (0, -1), (1, 1), (2, 3)]
w, b = Value(-1.5), Value(2.0)
parameters = [w, b]
state = [{"m": 0.0, "v": 0.0} for _ in parameters]
losses = []

for step in range(1, steps + 1):
zero_grad(parameters)
squared_errors = [(w * x + b - y) ** 2 for x, y in samples]
loss = sum(squared_errors) / len(samples)
loss.backward()
losses.append(loss.data)

if optimizer == "sgd":
sgd_step(parameters, learning_rate=0.1)
elif optimizer == "adam":
adam_step(parameters, state, step)
else:
raise ValueError("optimizer must be 'sgd' or 'adam'")

return {"w": w.data, "b": b.data,
"initial_loss": losses[0], "final_loss": losses[-1]}


if __name__ == "__main__":
manual_w, manual_b = Value(3), Value(-1)
manual_loss = (manual_w * 2 + manual_b - 4) ** 2 / 2
manual_loss.backward()
print(json.dumps({
"manual_grad": {"w": manual_w.grad, "b": manual_b.grad},
"sgd": train("sgd", 80),
"adam": train("adam", 220),
}, ensure_ascii=False))

运行结果应满足:手算例子的梯度为 w: 2,b: 1;两种优化器都把损失降到接近 0,并把参数学到 w ≈ 2,b ≈ -1。代码刻意把图的每个中间节点都重新创建,因此每轮训练结束后旧图可被释放。

这段实现省略了什么

教学实现具备 生产框架还要处理
标量加、乘、幂与梯度累积 张量形状、广播及批量高性能内核
动态建立本轮计算图 设备调度、混合精度与分布式通信
SGD、Adam 的核心更新 参数组、稀疏梯度、状态保存和融合算子
逆拓扑反向传播 高阶梯度、图释放策略、自定义算子与并发

10. 怎样确认梯度真的算对了

中心差分可以作为独立的数值检查:

1
∂L/∂θᵢ ≈ [L(θᵢ+h) - L(θᵢ-h)] / (2h)

它不依赖待测反向公式,适合检查自定义算子或教学实现。但 h 过大会产生截断误差,过小又会被浮点舍入淹没;ReLU 拐点等不可微位置也不适合直接比较。梯度检查是测试工具,不是训练算法。

本文的 tests/verify-backprop-article.cjs 会提取上面的 Python 代码,先检查手算结果和两种优化器的收敛,再生成 320 组表达式。JavaScript 参考端对每个变量做中心差分,Python 自动微分端做反向传播,共比较 960 个梯度值;两条路径没有共享导数实现。

11. 一个可靠训练步的顺序

以 PyTorch 风格的训练循环为例,顺序通常是:

顺序 操作 原因
1 optimizer.zero_grad() 清除上一轮累积梯度
2 prediction = model(batch) 建立本轮前向图
3 loss = criterion(prediction, target) 把训练目标归约为标量
4 loss.backward() 累积参数梯度
5 可选的梯度裁剪与诊断 在更新前观察或限制梯度
6 optimizer.step() 用梯度和优化器状态更新参数

若有意做梯度累积,可以连续处理多个微批次后再 step(),但通常要按累积步数缩放损失,并明确何时清零。自动混合精度还会涉及损失缩放、反缩放和溢出检测,不能只机械移动 step() 的位置。

12. 梯度消失、爆炸与裁剪

深层链式法则会连乘许多局部雅可比。若其典型尺度长期小于 1,早期层梯度可能逐层缩小;长期大于 1,则可能迅速放大。残差连接、归一化、合适初始化和激活函数都在不同角度改善信号与梯度传播。

梯度裁剪常按全局范数把过大的梯度缩回阈值内。它可以阻止一次异常更新破坏参数,却不会修复数据错误、不稳定损失、错误遮罩或持续不合理的学习率。

现象 先检查 不应立即下的结论
损失突然变成 NaN 输入范围、除零、对数定义域、混合精度溢出 一定只是学习率太大
梯度长期为 0 是否断图、饱和激活、冻结参数、遮罩与损失路径 模型已经收敛
梯度范数周期性尖峰 特定批次、序列长度、归一化和调度器 只要裁剪就彻底解决
训练损失降而验证变差 数据划分、过拟合、分布偏移 优化器没有工作

13. 最容易踩的十个坑

  1. 把反向传播等同于梯度下降。 前者求梯度,后者使用梯度。
  2. 忘记清零梯度。 多轮 .backward() 默认会向叶子 .grad 累加。
  3. 在前向中把张量变成普通数值。 不恰当的 detach、重新包装或跨库转换会切断图。
  4. 原地修改反向需要的值。 框架可能报版本不一致,或自定义实现悄悄算错。
  5. eval() 当成 no-grad。 模块模式与梯度记录互相独立。
  6. 只盯总损失。 还应记录学习率、梯度范数、参数范数和各损失分量。
  7. 看到梯度大就立刻裁剪。 先确认单位、批次归约方式和损失是否正确。
  8. 认为 Adam 不需要调学习率。 自适应缩放仍以全局学习率为基础。
  9. 对所有参数一律做权重衰减。 偏置和归一化参数常需单独分组判断。
  10. 用数值差分替代正确性证明。 它只能在抽样点发现不一致,还受不可微点与浮点误差影响。

14. 原始资料与下一步

资料核对日期:2026-09-20。

  1. PyTorch:Autograd mechanics:核对动态图、叶子梯度累积、保存张量、no-grad、inference mode 与 eval() 的区别。
  2. PyTorch:Automatic differentiation package:核对 backwardgrad、前向模式自动微分与梯度布局等当前接口。
  3. PyTorch:Optimizer 文档:核对优化器基本用法、逐样本梯度和性能相关接口。
  4. PyTorch:torch.optim.SGD:核对 momentum、dampening、weight decay 与 Nesterov 参数的实际定义。
  5. PyTorch:torch.optim.AdamW:核对解耦权重衰减、状态与参数组选项。
  6. Kingma 与 Ba:Adam:原始算法、矩估计和初始化偏差修正。

读完本文,可以回到 Transformer 入门,把注意力矩阵、前馈层和嵌入表都视为同一张可微计算图中的参数与中间量。下一步适合继续学习“神经网络训练诊断”,把这里的梯度、损失和更新规则连接到过拟合、验证集、学习率曲线与可观测性;也可以进入“微调与 LoRA”,理解为什么冻结大部分叶子参数后仍能沿剩余路径反向传播。