AdamW Optimizer:为什么要把权重衰减与梯度更新解耦
在上一篇 Adam Optimizer 文章中,我们知道 Adam 会同时维护梯度的一阶矩和二阶矩,并根据历史梯度为不同参数调整有效步长。
当模型出现过拟合时,我们还希望限制参数变得过大,于是经常会为优化器设置 weight_decay:
optimizer = torch.optim.AdamW(
model.parameters(),
lr=1e-3,
weight_decay=1e-2,
)
这里为什么常用 AdamW,而不是简单地给 Adam 加一个 L2 正则项?AdamW 中的字母 W 又代表什么?
一句话概括:
AdamW 把权重衰减从 Adam 的梯度与矩估计中分离出来,让“根据损失优化参数”和“按比例缩小参数”成为两件独立的事。
1. 先回顾 Adam
Adam 对当前梯度 $g_t$ 维护两种状态。
一阶矩:
$$
m_t=\beta_1m_{t-1}+(1-\beta_1)g_t
$$
二阶矩:
$$
v_t=\beta_2v_{t-1}+(1-\beta_2)g_t^2
$$
经过偏差修正后:
$$
\hat{m}_t=\frac{m_t}{1-\beta_1^t}
$$
$$
\hat{v}_t=\frac{v_t}{1-\beta_2^t}
$$
Adam 的参数更新为:
$$
\theta_t
=\theta_{t-1}
-\alpha\frac{\hat{m}_t}{\sqrt{\hat{v}_t}+\epsilon}
$$
其中:
- $\theta$ 是模型参数。
- $\alpha$ 是学习率。
- $m_t$ 记录近期梯度方向。
- $v_t$ 记录近期梯度平方的尺度。
如果对 Adam 本身还不熟悉,可以先阅读上一篇文章:
2. 为什么需要限制权重大小
神经网络参数很多,模型可能把训练数据中的细节甚至噪声也记住。
一种常见思路是:在保证拟合数据的同时,不希望权重无限增大。
较小的权重并不自动等于更好的模型,但限制权重规模通常可以:
- 减少模型对个别输入特征的过度依赖。
- 缓解过拟合。
- 让参数更新更加平稳。
- 改善部分任务上的验证集表现。
实现这一目标时,经常会遇到两个容易混淆的概念:
- L2 正则化。
- 权重衰减(Weight Decay)。
3. L2 正则化是什么
L2 正则化会在原始损失后面增加参数平方惩罚项:
$$
L_{total}(\theta)
=L_{data}(\theta)
+\frac{\lambda}{2}\lVert\theta\rVert_2^2
$$
其中:
- $L_{data}$ 是模型在训练数据上的原始损失。
- $\lambda$ 控制正则化强度。
- $\lVert\theta\rVert_2^2$ 是参数平方和。
对参数求导:
$$
\nabla_\theta L_{total}
=\nabla_\theta L_{data}+\lambda\theta
$$
如果把数据损失产生的梯度记为 $g_t$,加入 L2 后,送给优化器的梯度就变成:
$$
g_t^{L2}=g_t+\lambda\theta_{t-1}
$$
也就是说,L2 正则化通过“修改梯度”来影响参数。
4. 权重衰减是什么
权重衰减的直觉更加直接:每次更新时,都让参数按比例缩小一点。
如果暂时不考虑损失梯度,可以写成:
$$
\theta_t=(1-\alpha\lambda)\theta_{t-1}
$$
例如:
$$
\theta=10,\quad \alpha=0.1,\quad \lambda=0.01
$$
那么衰减一次后:
$$
\theta’=(1-0.1\times0.01)\times10=9.99
$$
参数每一步都按固定比例衰减,而不是把衰减项混入损失梯度。
5. 为什么在普通 SGD 中两者看起来一样
普通 SGD 加 L2 后:
$$
\theta_t
=\theta_{t-1}-\alpha(g_t+\lambda\theta_{t-1})
$$
展开:
$$
\theta_t
=(1-\alpha\lambda)\theta_{t-1}
-\alpha g_t
$$
这个式子恰好可以拆成两部分:
- $(1-\alpha\lambda)\theta_{t-1}$:缩小原参数。
- $-\alpha g_t$:根据损失梯度更新参数。
因此,对于不带自适应缩放的普通 SGD,在常见设置下,L2 正则化和权重衰减可以得到等价的更新形式。
这也导致很多代码和资料把二者混在一起讨论。
6. 为什么到了 Adam 中就不一样了
Adam 不会直接使用梯度 $g_t$,而是先用它更新 $m_t$ 和 $v_t$。
如果直接把 L2 项加入梯度:
$$
g_t’=g_t+\lambda\theta_{t-1}
$$
那么 Adam 实际记录的是:
$$
m_t=\beta_1m_{t-1}+(1-\beta_1)g_t’
$$
$$
v_t=\beta_2v_{t-1}+(1-\beta_2)(g_t’)^2
$$
这意味着 $\lambda\theta$ 不只是缩小参数,它还会:
- 进入一阶矩,影响历史更新方向。
- 进入二阶矩,影响自适应缩放尺度。
- 与真正的数据梯度混合。
- 对不同参数产生不同的有效衰减程度。
因此,在 Adam 这类自适应优化器中,“把 L2 项加入梯度”与“直接按比例衰减参数”不再等价。
7. AdamW 的解耦更新
AdamW 仍然使用原始数据梯度更新一阶矩和二阶矩:
$$
g_t=\nabla_\theta L_{data}
$$
$$
m_t=\beta_1m_{t-1}+(1-\beta_1)g_t
$$
$$
v_t=\beta_2v_{t-1}+(1-\beta_2)g_t^2
$$
然后将 Adam 更新和权重衰减分开执行:
$$
\theta_t
=\theta_{t-1}
-\alpha\frac{\hat{m}_t}{\sqrt{\hat{v}t}+\epsilon}
-\alpha\lambda\theta{t-1}
$$
也可以写成:
$$
\theta_t
=(1-\alpha\lambda)\theta_{t-1}
-\alpha\frac{\hat{m}_t}{\sqrt{\hat{v}_t}+\epsilon}
$$
这里的关键不是多减了一项,而是:
权重衰减项不进入 $m_t$ 和 $v_t$,因此不会污染 Adam 对梯度方向和梯度尺度的估计。
字母 W 就来自 Weight Decay。
8. 用“零梯度”理解解耦
假设某一步数据损失对参数的梯度刚好为 0:
$$
g_t=0
$$
Adam + L2
加入 L2 后,优化器看到的梯度仍然不是 0:
$$
g_t’=\lambda\theta
$$
它会进入一阶矩和二阶矩,并受到 Adam 自适应缩放的影响。
AdamW
AdamW 的梯度路径仍然看到 $g_t=0$,而权重衰减独立执行:
$$
\theta_t=(1-\alpha\lambda)\theta_{t-1}
$$
这种行为更容易解释:即使当前数据梯度为 0,参数仍按照设定比例衰减,但优化器的矩估计不会把衰减误认为数据梯度。
9. 手算一次 AdamW 更新
假设:
$$
\theta_0=1,\quad
g_1=2
$$
设置:
$$
\alpha=0.001,\quad
\lambda=0.01,\quad
\beta_1=0.9,\quad
\beta_2=0.999
$$
为了方便手算,暂时忽略非常小的 $\epsilon$。
第一步 Adam 的偏差修正结果为:
$$
\hat{m}_1=2
$$
$$
\hat{v}_1=4
$$
Adam 梯度更新量:
$$
\alpha\frac{\hat{m}_1}{\sqrt{\hat{v}_1}}
=0.001\times\frac{2}{2}
=0.001
$$
独立权重衰减量:
$$
\alpha\lambda\theta_0
=0.001\times0.01\times1
=0.00001
$$
最终参数:
$$
\theta_1
=1-0.001-0.00001
=0.99899
$$
可以看到,梯度更新和权重衰减各自有清晰的来源。
10. Adam、Adam + L2 和 AdamW 对比
| 方法 | 矩估计使用的梯度 | 参数是否直接衰减 | 特点 |
|---|---|---|---|
| Adam | $g_t$ | 否 | 没有额外正则化 |
| Adam + L2 | $g_t+\lambda\theta$ | 否 | L2 项也进入一阶矩和二阶矩 |
| AdamW | $g_t$ | 是 | 梯度更新与权重衰减解耦 |
如果 weight_decay=0,AdamW 不执行参数衰减,其核心梯度更新与 Adam 对齐。
AdamW 的优势并不是“比 Adam 多一个字母”,而是让正则化强度不再被自适应梯度缩放隐式改变。
11. PyTorch 中如何使用 AdamW
最常见的写法:
optimizer = torch.optim.AdamW(
model.parameters(),
lr=1e-3,
betas=(0.9, 0.999),
eps=1e-8,
weight_decay=1e-2,
)
当前 PyTorch 官方接口中,AdamW 的常见默认值包括:
| 参数 | 常见默认值 | 作用 |
|---|---|---|
lr |
1e-3 |
基础学习率 |
betas |
(0.9, 0.999) |
一阶矩和二阶矩的衰减系数 |
eps |
1e-8 |
数值稳定项 |
weight_decay |
1e-2 |
解耦式权重衰减强度 |
amsgrad |
False |
是否使用 AMSGrad 变体 |
即使库提供默认值,实践中也建议显式写出 lr 和 weight_decay,让实验配置更加清楚。
12. 使用 AdamW 完整训练一个 MLP
下面继续使用 XOR 二分类任务:
import torch
from torch import nn
torch.manual_seed(42)
X = torch.tensor([
[0.0, 0.0],
[0.0, 1.0],
[1.0, 0.0],
[1.0, 1.0],
])
y = torch.tensor([0, 1, 1, 0])
class XORNet(nn.Module):
def __init__(self):
super().__init__()
self.net = nn.Sequential(
nn.Linear(2, 8),
nn.Tanh(),
nn.Linear(8, 2),
)
def forward(self, x):
return self.net(x)
model = XORNet()
loss_fn = nn.CrossEntropyLoss()
optimizer = torch.optim.AdamW(
model.parameters(),
lr=1e-2,
weight_decay=1e-2,
)
for epoch in range(1000):
model.train()
logits = model(X)
loss = loss_fn(logits, y)
optimizer.zero_grad(set_to_none=True)
loss.backward()
optimizer.step()
if epoch % 200 == 0:
print(f"epoch={epoch}, loss={loss.item():.6f}")
model.eval()
with torch.no_grad():
logits = model(X)
predictions = logits.argmax(dim=1)
print("predictions:", predictions.tolist())
训练正常时,预测结果应接近:
predictions: [0, 1, 1, 0]
训练循环和 Adam 完全相同。AdamW 的解耦逻辑已经封装在 optimizer.step() 中,不需要手动修改损失函数。
13. 哪些参数通常不做权重衰减
一个常见实践是:对全连接层、卷积层等权重矩阵进行衰减,但对下面这些参数关闭衰减:
- 偏置
bias。 - LayerNorm、BatchNorm 等归一化层的缩放和偏移参数。
- 某些任务中特殊定义的标量参数。
原因主要是:
- 偏置和归一化参数通常维度较低。
- 它们的作用不同于大规模权重矩阵。
- 对它们衰减未必带来期望的正则化效果。
这不是不可违反的数学定律,而是一种常见实验配置。是否排除仍应以验证集结果为准。
14. 用参数组排除 bias 和归一化参数
可以根据参数名称进行分组:
decay_params = []
no_decay_params = []
for name, param in model.named_parameters():
if not param.requires_grad:
continue
is_bias = name.endswith("bias")
is_norm = "norm" in name.lower()
if is_bias or is_norm:
no_decay_params.append(param)
else:
decay_params.append(param)
optimizer = torch.optim.AdamW(
[
{
"params": decay_params,
"weight_decay": 1e-2,
},
{
"params": no_decay_params,
"weight_decay": 0.0,
},
],
lr=1e-3,
)
对于结构复杂的模型,只根据名称中是否包含 norm 进行判断可能不够严谨。
还可以使用一个简单经验:一维参数通常是偏置或归一化参数,二维及以上参数通常是权重矩阵。
decay_params = []
no_decay_params = []
for param in model.parameters():
if not param.requires_grad:
continue
if param.ndim >= 2:
decay_params.append(param)
else:
no_decay_params.append(param)
这种方法也不是对所有模型都绝对正确,但在 Transformer 风格模型中很常见。
15. 参数组中不能重复加入同一个参数
每个可训练参数只能出现在一个参数组里。
错误示例:
optimizer = torch.optim.AdamW([
{"params": model.parameters(), "weight_decay": 1e-2},
{"params": model.classifier.parameters(), "weight_decay": 0.0},
])
如果 classifier 本来就在 model.parameters() 中,它就被重复加入了。
正确做法是先把参数划分成互不重叠的集合,再创建优化器。
16. AdamW 与学习率调度
AdamW 仍然需要合理的学习率策略。
常见组合包括:
- AdamW + warmup。
- AdamW + cosine decay。
- AdamW + linear decay。
例如使用余弦退火:
optimizer = torch.optim.AdamW(
model.parameters(),
lr=1e-3,
weight_decay=1e-2,
)
scheduler = torch.optim.lr_scheduler.CosineAnnealingLR(
optimizer,
T_max=100,
)
for epoch in range(100):
train_one_epoch(model, optimizer)
scheduler.step()
因为 AdamW 的衰减量中包含学习率:
$$
\alpha\lambda\theta
$$
所以学习率随训练变化时,每一步实际施加的衰减量也会随之变化。
17. 如何选择 weight_decay
没有一个适合所有任务的固定值。
可以把下面的值作为搜索起点:
0
1e-4
1e-3
1e-2
1e-1
常见现象:
weight_decay太小:正则化效果不明显。weight_decay太大:参数被过度压缩,训练集都难以拟合。- 训练集很好、验证集较差:可以尝试适当增大。
- 训练集和验证集都较差:先检查模型容量、学习率和训练过程,不要只靠增大衰减解决。
lr 和 weight_decay 会相互影响,最好联合调节,而不是只搜索其中一个。
18. 保存与恢复 AdamW 状态
AdamW 和 Adam 一样,需要保存一阶矩、二阶矩和更新步数。
torch.save({
"model": model.state_dict(),
"optimizer": optimizer.state_dict(),
"scheduler": scheduler.state_dict(),
"epoch": epoch,
}, "checkpoint.pth")
恢复时:
checkpoint = torch.load("checkpoint.pth")
model.load_state_dict(checkpoint["model"])
optimizer.load_state_dict(checkpoint["optimizer"])
scheduler.load_state_dict(checkpoint["scheduler"])
start_epoch = checkpoint["epoch"] + 1
只恢复模型而不恢复优化器,AdamW 的历史矩估计会重新开始,训练轨迹可能发生变化。
19. 常见错误
19.1 同时手动加 L2 又设置 AdamW 衰减
例如既在损失中加入参数平方和:
loss = data_loss + lambda_ * l2_penalty
又设置:
optimizer = torch.optim.AdamW(
model.parameters(),
weight_decay=lambda_,
)
这会同时施加两种正则化。如果不是有意设计,通常属于重复计算。
19.2 认为 AdamW 一定优于 Adam
如果 weight_decay=0,两者没有解耦衰减上的差异。
如果任务本身不需要额外权重衰减,AdamW 不会凭空带来提升。
19.3 直接照搬别人的 weight_decay
合适的衰减强度与下面因素有关:
- 学习率。
- 训练步数。
- 模型规模。
- 数据量。
- 是否使用归一化层。
- 学习率调度方式。
所以应通过验证集或交叉验证选择,而不是机械照搬。
19.4 对所有参数一视同仁
简单任务可以直接传入 model.parameters(),但大型模型通常值得认真区分权重矩阵、偏置和归一化参数。
19.5 在训练循环中重新创建优化器
错误写法:
for epoch in range(100):
optimizer = torch.optim.AdamW(model.parameters())
这会不断丢失一阶矩、二阶矩和更新步数。优化器通常只在训练循环外创建一次。
20. AdamW 适合哪些场景
AdamW 经常作为这些任务的优化器起点:
- Transformer。
- BERT、GPT 等语言模型。
- Vision Transformer。
- 需要 Adam 自适应更新并希望使用权重衰减的任务。
- 预训练模型微调。
不过它并不限定于 Transformer。普通 MLP、CNN 或其他模型只要适合 Adam,并且需要解耦权重衰减,也可以使用 AdamW。
21. 总结
L2 正则化和权重衰减的目标看起来相似,都是限制参数规模,但它们的实现路径不同:
L2 正则化
-> 修改损失和梯度
-> 正则项进入 Adam 的一阶矩与二阶矩
AdamW
-> 使用原始数据梯度计算矩估计
-> 单独按比例衰减参数
最重要的结论是:
- 在普通 SGD 中,L2 正则化和权重衰减在常见设置下可以等价。
- 在 Adam 中,把 L2 项加入梯度会受到自适应缩放影响。
- AdamW 把权重衰减从梯度路径中解耦。
- AdamW 的
weight_decay仍然需要根据验证集调节。 - 大型模型中通常会对权重矩阵衰减,而排除偏置和归一化参数。
PyTorch 中的基础写法是:
optimizer = torch.optim.AdamW(
model.parameters(),
lr=1e-3,
weight_decay=1e-2,
)
这行代码虽然简单,但理解“解耦”以后,才真正知道它与 Adam + L2 的区别。
参考资料
- Ilya Loshchilov, Frank Hutter, Decoupled Weight Decay Regularization: https://arxiv.org/abs/1711.05101
- PyTorch Documentation,
torch.optim.AdamW: https://docs.pytorch.org/docs/stable/generated/torch.optim.adamw.AdamW_class.html - PyTorch Documentation,
torch.optim: https://docs.pytorch.org/docs/stable/optim.html - Diederik P. Kingma, Jimmy Ba, Adam: A Method for Stochastic Optimization: https://arxiv.org/abs/1412.6980
