AdamW:从梯度下降到解耦权重衰减
——结合ChatGPT整理
1. 最初只有参数、损失和梯度
设模型参数为
\[\theta \in \mathbb{R}^D.
\]
这里 \(\theta\) 不是一个数,而是把 Transformer 中所有待训练参数统一记成的参数向量,例如:
\[W_Q,\ W_K,\ W_V,\ W_O,\ W_{\mathrm{FFN}},\ \text{Embedding 等}.
\]
对于第 \(k\) 个 mini-batch,计算得到损失
\[\mathcal{L}_k(\theta).
\]
反向传播得到梯度:
\[\boxed{
g_k = \nabla_\theta \mathcal{L}_k(\theta_{k-1})
}
\]
其中:
- \(k\):当前是第几次参数更新;
- \(\theta_{k-1}\):更新前的参数;
- \(g_k\):当前 mini-batch 给出的梯度;
- \(g_{k,j}\):第 \(j\) 个参数对应的梯度。
梯度告诉我们:如果某个参数稍微增加,损失会怎样变化。
因此最简单的更新是梯度下降:
\[\boxed{
\theta_k = \theta_{k-1} - \eta g_k
}
\]
其中 \(\eta>0\) 是学习率。
减号的原因是:梯度指向损失增加最快的方向,所以负梯度指向局部下降最快的方向。
2. 普通梯度下降的第一个问题:梯度噪声很大
我们通常不会对整个训练集计算梯度,而只抽一个 mini-batch。
因此当前梯度
\[g_k = \nabla_\theta \mathcal{L}_k(\theta_{k-1})
\]
只是完整数据梯度的随机估计。
例如,连续几个 mini-batch 得到某个参数的梯度:
\[1.2,\quad -0.3,\quad 0.8,\quad 1.0,\quad -0.1.
\]
虽然总体方向可能是正的,但单个 batch 的梯度会来回波动。如果直接使用当前梯度更新:
\[\theta_k = \theta_{k-1} - \eta g_k,
\]
参数就会随着 mini-batch 噪声左右摇摆。
因此,我们不希望只相信当前的 \(g_k\),而希望综合过去若干步的梯度方向。
这就引出了 \(m_k\)。
3. 为什么引入 \(m_k\):估计稳定的梯度方向
定义:
\[\boxed{
m_k
=
\beta_1 m_{k-1}
+
(1-\beta_1)g_k
}
\]
其中:
- \(m_k\):截至第 \(k\) 步,梯度的一阶矩估计;
- \(m_{k-1}\):过去梯度信息的累计;
- \(g_k\):当前梯度;
- \(\beta_1\in[0,1)\):保留历史信息的比例。
这里的“一阶矩”可以暂时理解为平均值。更准确地说,\(m_k\) 是梯度的指数移动平均。
例如取
\[\beta_1 = 0.9,
\]
则
\[m_k = 0.9m_{k-1}+0.1g_k.
\]
意思是:
- 90% 来自过去积累的方向;
- 10% 来自当前梯度。
把递推展开:
\[\begin{aligned}
m_k
&=\beta_1m_{k-1}+(1-\beta_1)g_k\\
&=\beta_1\left[\beta_1m_{k-2}+(1-\beta_1)g_{k-1}\right]
+(1-\beta_1)g_k\\
&=\beta_1^2m_{k-2}
+(1-\beta_1)\beta_1g_{k-1}
+(1-\beta_1)g_k.
\end{aligned}
\]
如果初始化
\[m_0=0,
\]
继续展开得到
\[\boxed{
m_k
=
(1-\beta_1)
\sum_{i=1}^{k}
\beta_1^{k-i}g_i
}
\]
所以近期梯度权重大,久远梯度权重按照指数衰减。
例如:
\[m_k
=
(1-\beta_1)
\left(
g_k+\beta_1g_{k-1}+\beta_1^2g_{k-2}+\cdots
\right).
\]
因此 \(m_k\) 的作用是:
\[\boxed{
\text{平滑梯度噪声,保留一段时间内稳定的更新方向}
}
\]
如果只使用 \(m_k\),得到的就是带动量的优化方法:
\[\theta_k=\theta_{k-1}-\eta m_k.
\]
4. 只有 \(m_k\) 还不够:不同参数的梯度尺度差异很大
Transformer 中不同参数的梯度大小可能非常不同。
例如同一步中两个参数的梯度是:
\[g_{k,1}=0.001,
\qquad
g_{k,2}=10.
\]
如果所有参数都使用相同学习率 \(\eta\),更新量分别是
\[-\eta\cdot0.001,
\qquad
-\eta\cdot10.
\]
第二个参数的更新量是第一个参数的 \(10\,000\) 倍。
这会产生一个问题:
- 学习率设大,梯度较大的参数可能更新过猛甚至发散;
- 学习率设小,梯度较小的参数几乎不动。
因此我们希望为每个参数估计它“通常的梯度尺度”,然后做自适应归一化。
这就引出了 \(v_k\)。
5. 为什么引入 \(v_k\):估计每个参数的梯度尺度
定义:
\[\boxed{
v_k
=
\beta_2v_{k-1}
+
(1-\beta_2)g_k^2
}
\]
这里所有运算都是逐元素进行的。
如果
\[g_k=
\begin{bmatrix}
g_{k,1}\\
g_{k,2}
\end{bmatrix},
\]
那么
\[g_k^2=
\begin{bmatrix}
g_{k,1}^2\\
g_{k,2}^2
\end{bmatrix}.
\]
这不是向量的平方范数,也不是矩阵乘法。
\(v_k\) 的含义是:过去梯度平方的指数移动平均。它被称为梯度的二阶原点矩估计。
需要特别注意:
\[\boxed{
v_k\text{ 不是严格意义上的方差}
}
\]
因为方差应当是
\[\operatorname{Var}(g)
=
\mathbb{E}[g^2]-\mathbb{E}[g]^2,
\]
而 \(v_k\) 只估计其中的
\[\mathbb{E}[g^2].
\]
Adam 论文里常把它简称为 second moment,即二阶矩。
5.1 为什么要平方?
假设某个参数最近的梯度为
\[10,\quad -10,\quad 10,\quad -10.
\]
如果直接计算平均值,可能得到接近零:
\[\frac{10-10+10-10}{4}=0.
\]
但这并不意味着它的梯度尺度小。恰恰相反,它的梯度振幅很大。
平方以后:
\[100,\quad100,\quad100,\quad100,
\]
便能正确反映该参数的梯度很大。
因此平方有两个作用:
- 消除正负号,避免不同方向抵消;
- 衡量梯度的大小或能量。
5.2 为什么不只用当前的 \(g_k^2\)?
如果直接用
\[v_k=g_k^2,
\]
那么归一化后会出现:
\[\frac{g_k}{\sqrt{g_k^2}}
=
\frac{g_k}{|g_k|}
=
\operatorname{sign}(g_k).
\]
这样梯度大小信息几乎全部丢失,而且某个 mini-batch 的异常梯度会立即改变步长,非常不稳定。
因此需要对历史梯度平方做平滑:
\[v_k
=
\beta_2v_{k-1}
+
(1-\beta_2)g_k^2.
\]
把它展开:
\[\boxed{
v_k
=
(1-\beta_2)
\sum_{i=1}^{k}
\beta_2^{k-i}g_i^2
}
\]
通常取:
\[\beta_2=0.999.
\]
这意味着 \(v_k\) 会在较长时间尺度上估计梯度大小,而不会因为某一个 batch 突然改变。
6. 为什么使用 \(\sqrt{v_k}\),而不是直接使用 \(v_k\)?
\(v_k\) 估计的是梯度平方:
\[v_k\approx \mathbb{E}[g^2].
\]
它的量纲相当于“梯度的平方”。
为了恢复到和梯度相同的尺度,要取平方根:
\[\sqrt{v_k}\approx \sqrt{\mathbb{E}[g^2]}.
\]
这就是梯度的均方根尺度,英文为 root mean square,简称 RMS。
于是可以把更新方向写成:
\[\frac{m_k}{\sqrt{v_k}}.
\]
直观上:
- 若某个参数长期梯度很大,则 \(v_{k,j}\) 大,分母大,更新被压小;
- 若某个参数长期梯度较小,则 \(v_{k,j}\) 小,分母小,相对更新被放大。
所以每个参数实际上拥有自己的有效学习率:
\[\boxed{
\eta_{k,j}^{\mathrm{effective}}
=
\frac{\eta}{\sqrt{v_{k,j}}}
}
\]
这就是 Adam 中 “adaptive”,即自适应的来源。
7. 为什么同时需要 \(m_k\) 和 \(v_k\)?
二者负责不同的问题。
\(m_k\) 解决的是:
\[\boxed{
\text{往哪个方向更新更可靠}
}
\]
它平滑梯度的正负方向。
\(v_k\) 解决的是:
\[\boxed{
\text{这个参数应该走多大一步}
}
\]
它估计梯度尺度,用来调节步长。
因此可以先得到一个初步更新:
\[\theta_k
=
\theta_{k-1}
-
\eta
\frac{m_k}{\sqrt{v_k}}.
\]
但是这个公式在训练刚开始时还有偏差问题。
8. 为什么需要偏差修正?
Adam 通常初始化:
\[m_0=0,
\qquad
v_0=0.
\]
考虑第一步:
\[m_1=(1-\beta_1)g_1.
\]
如果
\[\beta_1=0.9,
\]
那么
\[m_1=0.1g_1.
\]
它明显小于当前梯度 \(g_1\),原因不是梯度真的小,而是历史状态从零开始。
同理:
\[v_1=(1-\beta_2)g_1^2.
\]
如果
\[\beta_2=0.999,
\]
则
\[v_1=0.001g_1^2.
\]
也严重偏小。
这种偏小来自初始化为零,因此称为初始化偏差。
8.1 \(m_k\) 的偏差因子是怎么来的?
假设一段时间内梯度的期望近似不变:
\[\mathbb{E}[g_i]=\mu.
\]
由
\[m_k
=
(1-\beta_1)
\sum_{i=1}^{k}\beta_1^{k-i}g_i
\]
取期望:
\[\begin{aligned}
\mathbb{E}[m_k]
&=
(1-\beta_1)
\sum_{i=1}^{k}\beta_1^{k-i}\mathbb{E}[g_i]\\
&=
(1-\beta_1)
\sum_{i=1}^{k}\beta_1^{k-i}\mu.
\end{aligned}
\]
利用等比数列:
\[\sum_{i=1}^{k}\beta_1^{k-i}
=
1+\beta_1+\cdots+\beta_1^{k-1}
=
\frac{1-\beta_1^k}{1-\beta_1}.
\]
所以:
\[\begin{aligned}
\mathbb{E}[m_k]
&=
(1-\beta_1)
\frac{1-\beta_1^k}{1-\beta_1}\mu\\
&=
(1-\beta_1^k)\mu.
\end{aligned}
\]
我们真正想估计的是 \(\mu\),但 \(m_k\) 只估计到
\[(1-\beta_1^k)\mu.
\]
因此除以这个缺失因子:
\[\boxed{
\hat m_k
=
\frac{m_k}{1-\beta_1^k}
}
\]
这样近似有:
\[\mathbb{E}[\hat m_k]\approx\mu.
\]
8.2 \(v_k\) 的偏差修正同理
假设:
\[\mathbb{E}[g_i^2]=\nu.
\]
则可推出:
\[\mathbb{E}[v_k]
=
(1-\beta_2^k)\nu.
\]
因此定义:
\[\boxed{
\hat v_k
=
\frac{v_k}{1-\beta_2^k}
}
\]
用于修正训练初期 由于零初始化造成的低估。
随着 \(k\) 增大:
\[\beta_1^k\to0,
\qquad
\beta_2^k\to0,
\]
偏差修正逐渐接近 1,影响自然消失。
9. 为什么分母还要加入 \(\epsilon\)?
现在更新式为:
\[\theta_k
=
\theta_{k-1}
-
\eta
\frac{\hat m_k}{\sqrt{\hat v_k}}.
\]
但某些参数的梯度可能一直是零,于是:
\[\hat v_{k,j}=0.
\]
此时分母为零,无法计算。
即使 \(\hat v_{k,j}\) 不是严格等于零,而只是非常小,也可能导致:
\[\frac{1}{\sqrt{\hat v_{k,j}}}
\]
异常大,引起数值不稳定。
因此在分母中加入一个很小的正数:
\[\boxed{
\epsilon>0
}
\]
得到:
\[\boxed{
\theta_k
=
\theta_{k-1}
-
\eta
\frac{\hat m_k}
{\sqrt{\hat v_k}+\epsilon}
}
\]
常见取值是:
\[\epsilon=10^{-8}.
\]
它的主要作用是防止除零和提高数值稳定性,不是主要的学习率调节项。
至此得到的就是 Adam。
10. Adam 的完整逻辑
第 \(k\) 步先计算梯度:
\[g_k
=
\nabla_\theta\mathcal{L}_k(\theta_{k-1}).
\]
为了减少 mini-batch 梯度方向的噪声,引入一阶矩:
\[m_k
=
\beta_1m_{k-1}
+
(1-\beta_1)g_k.
\]
为了估计每个参数的梯度尺度,引入二阶矩:
\[v_k
=
\beta_2v_{k-1}
+
(1-\beta_2)g_k^2.
\]
为了修正零初始化导致的低估:
\[\hat m_k
=
\frac{m_k}{1-\beta_1^k},
\qquad
\hat v_k
=
\frac{v_k}{1-\beta_2^k}.
\]
然后更新:
\[\boxed{
\theta_k
=
\theta_{k-1}
-
\eta
\frac{\hat m_k}
{\sqrt{\hat v_k}+\epsilon}
}
\]
其中:
- \(\hat m_k\):平滑后的可靠方向;
- \(\sqrt{\hat v_k}\):该参数正常的梯度尺度;
- 二者相除:得到经过尺度归一化的方向;
- \(\eta\):统一控制总体更新速度;
- \(\epsilon\):防止数值问题。
11. 为什么还要从 Adam 变成 AdamW?
Adam 已经解决了方向和尺度问题,但训练大模型时还常希望限制权重无限增大。
一种常见方法是向损失加入 \(L_2\) 正则化:
\[\mathcal{L}_{\mathrm{reg}}(\theta)
=
\mathcal{L}(\theta)
+
\frac{\lambda}{2}\lVert\theta\rVert_2^2.
\]
这里:
- \(\lVert\theta\rVert_2^2\):所有参数平方之和;
- \(\lambda\):正则化强度。
对它求梯度:
\[\nabla_\theta\mathcal{L}_{\mathrm{reg}}
=
\nabla_\theta\mathcal{L}
+
\lambda\theta.
\]
也就是把原梯度
\[g_k
\]
改成:
\[g_k+\lambda\theta_{k-1}.
\]
12. 为什么在 Adam 中不能简单地把 \(L_2\) 正则项混入梯度?
如果直接把
\[g_k+\lambda\theta_{k-1}
\]
送入 Adam,那么 \(\lambda\theta\) 也会进入 \(m_k\)、\(v_k\),最后被自适应分母缩放。
粗略看,参数 \(j\) 的正则更新会变成类似:
\[-\eta
\frac{\lambda\theta_j}
{\sqrt{\hat v_{k,j}}+\epsilon}.
\]
这意味着:
- \(\hat v_{k,j}\) 大的参数,衰减较弱;
- \(\hat v_{k,j}\) 小的参数,衰减较强。
于是原本统一的正则化强度 \(\lambda\),被 Adam 的自适应机制扭曲了。
在普通 SGD 中,\(L_2\) 正则和权重衰减基本等价;但在 Adam 这种自适应优化器中,二者不再等价。
所以 AdamW 的做法是:
不把权重衰减混入梯度,而是把它作为独立的参数缩小步骤。
13. 为什么权重衰减写成 \(-\eta\lambda\theta\)?
我们希望每一步让参数按一定比例缩小:
\[\theta
\leftarrow
(1-\eta\lambda)\theta.
\]
展开:
\[\theta
\leftarrow
\theta-\eta\lambda\theta.
\]
所以权重衰减项是:
\[-\eta\lambda\theta.
\]
其中:
- \(\lambda\):规定衰减强度;
- \(\eta\):当前步长;
- \(\theta\):参数越大,衰减绝对值越大。
它不是每步减去固定常数,而是按参数当前大小做比例收缩。
14. 最终得到 AdamW
Adam 的梯度更新部分是:
\[-\eta
\frac{\hat m_k}
{\sqrt{\hat v_k}+\epsilon}.
\]
独立的权重衰减部分是:
\[-\eta\lambda\theta_{k-1}.
\]
合起来:
\[\boxed{
\theta_k
=
\theta_{k-1}
-
\eta
\frac{\hat m_k}
{\sqrt{\hat v_k}+\epsilon}
-
\eta\lambda\theta_{k-1}
}
\]
也可以写成:
\[\boxed{
\theta_k
=
(1-\eta\lambda)\theta_{k-1}
-
\eta
\frac{\hat m_k}
{\sqrt{\hat v_k}+\epsilon}
}
\]
这就是 AdamW。
15. 最后把每个量的来源串起来
\[g_k=\nabla_\theta\mathcal{L}_k
\]
来自反向传播,表示当前 mini-batch 建议的更新方向。
但 \(g_k\) 噪声较大,所以引入:
\[m_k
\]
对历史梯度求指数移动平均,估计稳定方向。
不同参数梯度尺度不同,所以引入:
\[v_k
\]
对历史梯度平方求指数移动平均,估计每个参数的梯度尺度。
因为 \(m_0=v_0=0\),训练初期两者偏小,所以引入:
\[\hat m_k,\qquad \hat v_k
\]
进行偏差修正。
因为 \(\sqrt{\hat v_k}\) 可能为零或过小,所以引入:
\[\epsilon
\]
保证数值稳定。
为了控制整体步长,引入:
\[\eta
\]
作为学习率。
为了让模型参数不要无限变大、改善正则化效果,引入:
\[\lambda
\]
作为权重衰减系数,并且不把它混入 Adam 梯度,而是单独衰减参数。
因此整个 AdamW 不是凭空设计出来的一条复杂公式,而是连续解决这些问题得到的:
\[\boxed{
\begin{aligned}
&\text{梯度有噪声}
&&\Rightarrow m_k,\\
&\text{各参数梯度尺度不同}
&&\Rightarrow v_k,\\
&\text{零初始化产生偏差}
&&\Rightarrow \hat m_k,\hat v_k,\\
&\text{分母可能过小}
&&\Rightarrow \epsilon,\\
&\text{需要控制总体步长}
&&\Rightarrow \eta,\\
&\text{需要独立正则化参数}
&&\Rightarrow \lambda\text{ 和 decoupled weight decay}.
\end{aligned}
}
\]
在 LLaDA 中,带 \(1/t\) 权重的 Mask Token 交叉熵负责产生梯度 \(g_k\);AdamW 并不关心梯度是怎样由扩散目标得到的,它只负责依据 \(g_k\)、历史一阶矩、历史二阶矩和权重衰减来更新 Transformer 参数。