梯度方向反复横跳怎么办?Momentum 与 Adam 如何重塑更新步长
从狭长谷底中的 SGD 振荡出发,手算 Momentum 与 Adam 的状态更新、偏差修正和逐参数缩放,并用 PyTorch 2.13 构建可调试训练循环。
上一篇用 Xavier/He 初始化守住了训练起点的信号尺度。现在反向传播能给出每个参数的梯度 ,但“知道当前位置最陡的下坡方向”不等于“能快速走到谷底”:在狭长曲面中,普通随机梯度下降(Stochastic Gradient Descent,SGD)会横向来回摆动,纵向却进展缓慢。
本文只研究一个核心问题:优化器怎样把当前梯度和历史状态组合成真正的参数更新? 我们从 Momentum 的方向平滑走到 Adam 的逐参数尺度适配,手算状态、写出张量数据流,并明确 PyTorch 2.13 中最容易被忽略的实现语义。
01 同一个学习率为什么顾不过来两个方向?#
考虑二维二次目标:
梯度为:
方向曲率很大,稍微偏离就产生大梯度; 方向平缓,梯度很小。SGD 更新:
若学习率 足够大以快速推进 , 可能越过谷底甚至发散;若把 降到稳定, 又移动得很慢。
等高线中的更新轨迹(示意)
陡峭 x 方向 ◄────────►
╲ ╱
╲ SGD╱ 左右梯度交替,更新抵消
╲ ╱
│
│ 平缓 y 方向:真正希望持续前进
▼
Momentum:削弱反复变号的横向分量,积累方向一致的纵向分量
Adam:再按每个参数近期梯度平方的尺度归一化更新textmini-batch 噪声还会让 抖动。我们需要的不是抛弃梯度,而是为每个参数保存少量历史状态,将短期噪声与长期方向分开。
02 Momentum 怎样积累“速度”?#
动量法(Momentum)维护与参数同形的缓冲 。一种常见写法是:
其中:
- :第 步后的全部参数;
- :当前 mini-batch 梯度;
- :动量缓冲,与参数逐元素对应;
- :学习率;:动量系数,常见起点是 0.9。
展开递推:
越早的梯度按指数衰减。若某方向的梯度一直同号,贡献会累积;若正负交替,贡献会互相抵消。这正好对应狭长谷底中的“纵向加速、横向减振”。
03 用两维梯度手算三步 Momentum#
设 ,三步梯度为:
第一个分量反复变号,第二个始终为正。
第 1 步:
第 2 步:
第 3 步:
三步原始梯度求和为 ;动量缓冲末值为 。第一个方向因反复变号没有无界积累,第二个方向从 1 增至 2.71。注意 Momentum 并不知道哪个方向是“正确的”,它只利用了梯度方向的时间一致性。
04 Nesterov 为什么要在“将要到达的位置”看梯度?#
Nesterov 动量(Nesterov Accelerated Gradient,NAG)的思想是先按历史速度向前看,再在预估位置计算梯度。不同教材和框架会使用代数等价或尺度不同的缓冲定义,因此代码审查时不能只凭变量名 velocity 判断公式。
概念形式可写为:
“提前看”可以更早纠正高速越过谷底的趋势。但 Nesterov 不是免费提速开关;学习率和动量仍需一起验证,而且应以所用框架的官方算法说明为准。
05 Adam 为什么还要记录梯度平方?#
Momentum 对所有参数使用同一全局学习率。Adam(Adaptive Moment Estimation)再维护梯度的一阶矩与二阶原始矩指数平均:
平方是逐元素的。 平滑方向, 估计每个参数近期梯度平方尺度。最终更新为:
其中 都与参数同形。梯度长期较大的参数分母也大,单步会被缩小;稀疏或尺度较小的方向可能得到相对更大的有效步长。
这不是近似 Hessian 的完整二阶优化。Adam 只使用逐坐标的梯度平方,没有保存参数间的曲率耦合。
06 为什么必须做偏差修正?#
。训练早期,指数平均会因为从 0 启动而偏小。Adam 使用:
用一个常梯度 的标量例子,取 。
第 1 步:
修正前两者明显小于真实一、二阶矩;修正后:
忽略很小的 ,第 1 步更新量为:
若漏掉偏差修正,早期有效步长会被错误地改变,尤其 很接近 1 时更明显。
07 两个参数尺度相差百倍时会怎样?#
设同一步梯度为 ,。第 1 步偏差修正后:
于是:
Adam 的首次更新几乎只保留符号,两个参数都走约一个 。这解释了它对梯度尺度差异的适应性,也揭示一个限制:参数真实需要的函数空间步长未必应该相同;逐坐标归一化可能改变隐含优化偏好。
08 从梯度到更新的完整数据流#
X [N,D] ─► model(θ) ─► prediction [N,C] ─► loss []
│
▼ backward
gradient g_t,形状同 θ
│
┌────────────────────┴───────────────────┐
▼ ▼
Momentum: v_t = μv_{t-1}+g_t Adam: m_t, v_t, step=t
│ │ 偏差修正 + 逐元素除法
└────────────────────┬───────────────────┘
▼
update Δθ,形状同 θ
│
▼
θ ← θ - Δθ
optimizer state 不是梯度:
参数 θ [P];梯度 g [P];Momentum 额外约 [P];Adam 额外约 [2P] + steptext以十亿参数模型为例,仅 Adam 的两个同精度状态就约等于额外二十亿个数,还未计参数、梯度、主权重副本和激活。这是选择优化器时真实的显存/内存成本。
09 不依赖优化器黑盒,写出最小 Adam#
import numpy as np
def adam_step(parameter, gradient, state, *, lr=1e-3,
beta1=0.9, beta2=0.999, eps=1e-8):
"""所有数组形状相同;返回新参数与新状态。"""
step = state["step"] + 1
first = beta1 * state["first"] + (1.0 - beta1) * gradient
second = beta2 * state["second"] + (1.0 - beta2) * gradient**2
first_hat = first / (1.0 - beta1**step)
second_hat = second / (1.0 - beta2**step)
update = lr * first_hat / (np.sqrt(second_hat) + eps)
new_parameter = parameter - update
new_state = {"step": step, "first": first, "second": second}
return new_parameter, new_state
parameter = np.array([1.0, -1.0])
state = {
"step": 0,
"first": np.zeros_like(parameter),
"second": np.zeros_like(parameter),
}
parameter, state = adam_step(
parameter,
gradient=np.array([0.1, 10.0]),
state=state,
lr=0.01,
)
np.testing.assert_allclose(parameter, [0.99, -1.01], rtol=0, atol=1e-7)python这个实现刻意没有权重衰减、AMSGrad、稀疏梯度和混合精度分支,目的是让每个状态可手查。生产代码应使用经过测试的官方优化器,但先理解状态转移,才能解释 checkpoint、恢复训练与显存占用。
10 用 PyTorch 2.13 正确落地#
当前 PyTorch 2.13 官方 torch.optim.SGD ↗ 与 torch.optim.Adam ↗ 都接收参数迭代器,并在 step() 时读取参数的 .grad。
import torch
from torch import nn
model = nn.Sequential(
nn.Linear(64, 128),
nn.ReLU(),
nn.Linear(128, 1),
)
loss_fn = nn.BCEWithLogitsLoss()
# 二选一,而不是同时更新同一批参数
optimizer = torch.optim.Adam(
model.parameters(),
lr=3e-4,
betas=(0.9, 0.999),
eps=1e-8,
weight_decay=0.0,
)
model.train()
for features, targets in train_loader:
# features [N,64] float;targets [N] float,值为 0/1
optimizer.zero_grad(set_to_none=True)
logits = model(features).squeeze(-1) # [N]
loss = loss_fn(logits, targets) # []
loss.backward() # 写入 parameter.grad
grad_norm = torch.nn.utils.clip_grad_norm_(
model.parameters(), max_norm=5.0,
)
if not torch.isfinite(grad_norm):
raise FloatingPointError("non-finite gradient")
optimizer.step() # 更新参数与 optimizer.statepython若选择 Momentum SGD:
optimizer = torch.optim.SGD(
model.parameters(),
lr=0.05,
momentum=0.9,
dampening=0.0,
nesterov=True,
)python官方文档指出,PyTorch SGD 的动量缓冲在第一步初始化为当前梯度,而不是全 0;因此第一步动量不受 dampening 缩放,dampening 从第二步开始生效。Nesterov 还要求非零 momentum,并需满足该 API 的参数约束。复现实验时应记录框架、版本和完整优化器参数,不能只写“用了 Momentum”。
11 参数组怎样表达“同一模型,不同学习率”?#
参数组(Parameter Group)允许给不同参数设置不同超参数,例如对预训练骨干使用更小学习率:
optimizer = torch.optim.AdamW(
[
{"params": backbone.parameters(), "lr": 1e-5},
{"params": head.parameters(), "lr": 3e-4},
],
betas=(0.9, 0.999),
weight_decay=0.01,
)python每个参数只能出现在一个参数组。构造后应检查:
seen = set()
for group_index, group in enumerate(optimizer.param_groups):
print(group_index, group["lr"], group["weight_decay"])
for parameter in group["params"]:
assert id(parameter) not in seen, "parameter appears twice"
seen.add(id(parameter))pythonAdamW 使用解耦权重衰减(Decoupled Weight Decay):衰减不先混入 Adam 的一、二阶矩。它与把 加进梯度的 L2 惩罚,在自适应优化器中并不等价。偏置和归一化参数是否衰减应由模型与实验决定,不应机械套用。
12 保存模型时为什么还必须保存优化器?#
Momentum 的 、Adam 的 和步数 都会影响下一步更新。只恢复参数 model.state_dict(),却新建空优化器,相当于中途清空速度、二阶矩和偏差修正时钟。
checkpoint = {
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"epoch": epoch,
"global_step": global_step,
}
torch.save(checkpoint, "checkpoint.pt")
# 恢复时先构造相同模型与优化器,再加载状态
checkpoint = torch.load("checkpoint.pt", map_location="cpu", weights_only=True)
model.load_state_dict(checkpoint["model"])
optimizer.load_state_dict(checkpoint["optimizer"])python若还有学习率调度器(Learning-rate Scheduler)和混合精度 scaler,也要一起保存。加载后打印每个参数组学习率,并用连续小数据对比“不中断训练”和“保存后恢复”的下一步结果。
13 怎样观察优化器到底做了什么?#
只看 loss 曲线,无法区分“梯度太小”“学习率太小”和“Adam 分母太大”。每隔一段步数记录:
- 全局与逐层参数范数 、梯度范数 ;
- 相对更新比 ;
- 当前真实学习率(调度后),而不是配置初值;
- Momentum 缓冲范数,或 Adam 的
exp_avg、exp_avg_sq分布; - 梯度非有限、裁剪触发频率与被裁剪前的范数;
- 训练 loss、验证指标和 wall-clock/每步耗时。
可在一次 step() 前后做差:
before = {
name: parameter.detach().clone()
for name, parameter in model.named_parameters()
}
optimizer.step()
for name, parameter in model.named_parameters():
delta = parameter.detach() - before[name]
relative = delta.norm() / (before[name].norm() + 1e-12)
print(name, "update_norm=", delta.norm().item(),
"relative_update=", relative.item())python该方法会复制参数,只适合短期诊断。大模型可按层采样或在优化器状态中读取统计,避免每步翻倍显存。
14 常见错误与最短调试路径#
- 忘记
zero_grad。 新梯度会累加,优化器看到的是多批之和;若确实做梯度累积,应按累积步数缩放 loss,并只在边界step()。 - 先
step()后backward()。 此时没有当前梯度,参数不会按预期更新。 - 训练中途无意重建 optimizer。 Adam/Momentum 状态被清空;检查
global_step和optimizer.state大小。 - 只调 betas,不先扫学习率。 学习率通常是一阶敏感项;先在合理范围做短跑,再精调动量与衰减。
- 把梯度裁剪当学习率调度。 高频裁剪会改变更新方向和尺度;记录触发率并修复发散根因。
- Adam 的
eps在低精度中过小。 先确认状态张量 dtype、混合精度策略和非有限值来源,再依据官方实现与硬件调节。 - 把 weight decay 当成完全等价的 L2。 对 Adam 应明确使用耦合还是解耦形式,并记录实现。
- 恢复 checkpoint 后学习率错位。 优化器和 scheduler 的加载顺序、参数组结构必须与保存时一致;恢复后立即打印核对。
最短调试路径:固定一个小 batch → 关闭随机数据增强 → 验证 loss 与梯度有限 → 检查参数确实变化 → 打印真实学习率与更新比 → 比较 SGD、Momentum、Adam 的短轨迹,而不是直接跑完整实验。
15 它们会在哪些场景失败?#
- Momentum 在梯度方向长期错误或学习率过大时会带着惯性冲得更远;
- Adam 对超参数更宽容不等于无需调参,也不保证验证集泛化优于 SGD;
- 稀疏参数、嵌入表和超大模型可能受优化器状态内存限制,需要专用稀疏或分片方案;
- 强噪声、非平稳目标会让历史矩过时, 太大时适应变慢;
- 逐坐标自适应依赖参数化方式,重参数化后轨迹可能明显改变;
- 优化训练损失更快,不代表解决数据泄漏、标签噪声、分布偏移或过拟合。
| 方法 | 保存状态 | 每步核心 | 典型优势 | 主要代价/边界 |
|---|---|---|---|---|
| SGD | 无 | 当前梯度 | 简单、省内存、基线清晰 | 狭长谷底易振荡 |
| Momentum | 一阶缓冲 | 历史方向平滑 | 抑制交替方向,持续方向加速 | 多一份参数级状态 |
| RMSProp | 二阶平方平均 | 逐参数尺度归一化 | 适应梯度尺度 | 不含 Adam 式一阶矩组合 |
| Adam | 一阶 + 二阶 + 步数 | 平滑方向并自适应缩放 | 常见任务起步稳、调试友好 | 约两份状态,泛化并非总优 |
| AdamW | 同 Adam | Adam + 解耦衰减 | 衰减语义更清晰 | 衰减率仍需验证 |
16 今天真正需要记住什么?#
- SGD 的单一学习率在不同曲率方向间会冲突,mini-batch 噪声又会放大轨迹抖动。
- Momentum 累积方向一致的梯度、抵消反复变号的梯度;缓冲与参数同形。
- Adam 用一阶矩平滑方向、二阶原始矩缩放每个参数,并用 修正零初始化偏差。
- 优化器是有状态算法;checkpoint 若不保存 optimizer state,就没有真正连续训练。
- 选择优化器要同时看验证表现、更新比、状态内存与每步吞吐,不能只比较前几百步训练 loss。
17 思考题与小练习#
- 对梯度序列 ,手算 时三步 Momentum 缓冲和参数变化。再把梯度全改为 2,比较最终速度。
- 取 Adam 的 、,手算每一步 和单位学习率更新量。
- 在同一微型二次问题上分别运行 SGD、Momentum 和 Adam;记录每步参数、梯度、更新比与状态。保存第 20 步 checkpoint,恢复后验证第 21 步与不中断运行完全一致。
相关工作#
- Polyak (1964), Some Methods of Speeding Up the Convergence of Iteration Methods ↗:重球动量方法的经典来源。
- Nesterov (1983), A Method for Solving the Convex Programming Problem with Convergence Rate O(1/k²) ↗:加速梯度方法的早期工作。
- Duchi, Hazan & Singer (2011), Adaptive Subgradient Methods for Online Learning and Stochastic Optimization ↗:AdaGrad 逐坐标自适应学习率的代表性论文。
- Kingma & Ba (2015), Adam: A Method for Stochastic Optimization ↗:Adam 的一、二阶矩估计与偏差修正。
- Loshchilov & Hutter (2019), Decoupled Weight Decay Regularization ↗:阐明自适应优化器中 L2 惩罚与解耦权重衰减的差异。
18 下一篇预告#
初始化只控制训练起点,Momentum/Adam 只重塑参数更新;随着权重变化,中间激活的尺度仍会漂移。下一篇将比较 Batch Normalization 与 Layer Normalization 的统计轴、训练/推理数据流和适用架构,解释归一化为何不是“把所有张量都标准化”。