从零手写大模型 · 训练篇 04 - Adam Optimizer(Adam 优化器)

上期回顾

在训练篇 EP03里,我们把梯度下降(Gradient Descent)的更新规则从"感觉上应该往梯度反方向走"落地成了一行公式:

θ←θ−η⋅∇θL\theta \leftarrow \theta - \eta \cdot \nabla_\theta L

然后沿用"猫吃鱼"这个2层小例子,手动做了一步更新(学习率 η=0.1\eta = 0.1),loss从约 0.1914 降到了约 0.0157。我们也顺手过了一遍batch / mini-batch / SGD的区别。

但如果你把这个更新规则原封不动地套到真实训练里,很快会撞见两个很现实的问题。这就是今天的起点。

一、朴素梯度下降的两个痛点

痛点一:方向会抖。

想象损失函数的地形不是一个规规矩矩的碗,而是一条窄窄的峡谷——某个方向陡,某个方向缓。朴素梯度下降在陡的方向上会来回震荡(每步都冲过头,再被拉回来),在缓的方向上却挪得特别慢。整个下降路径像蛇形一样左右摆动,而不是径直冲向谷底。

痛点二:所有参数共用一个学习率。

在一个几十万甚至上亿参数的模型里,不同参数的梯度量级可能差好几个数量级。有的参数梯度一直很大,学习率稍微设大一点就会震荡发散;有的参数梯度一直很小,同样的学习率下几乎不动。用同一个 η\eta 去伺候所有参数,注定是"一刀切"式的妥协。

这两个痛点分别对应两类改进思路:动量(Momentum)解决"抖",自适应学习率(如RMSProp)解决"一刀切"。Adam把这两者结合在了一起,这也是它至今仍是深度学习默认优化器的原因。

二、Momentum:给梯度下降加上惯性

2.1 直觉

朴素梯度下降每一步只看"当前这一刻"的梯度,完全没有记忆。假设某个参数在训练中,梯度一开始是正的,然后慢慢减小、越过零点、变成负的:

g: +0.4→+0.2→0→−0.2→−0.4g:\ +0.4 \to +0.2 \to 0 \to -0.2 \to -0.4

如果没有动量,梯度一变负,参数立刻就要往反方向走——这就容易在最低点附近来回过头、来回震荡。Momentum的想法很简单:让更新量带上一点"惯性",不要只听当前这一步梯度的,也要参考之前一直积累的趋势。

物理类比很贴切:一个小球从山坡滚下来,如果只受"当前位置的坡度"支配,遇到小坑洼就会卡住来回晃;但如果小球有质量、有惯性,它会在整体下坡方向上越滚越快,局部的小震荡会被惯性"平均"掉,方向也不会因为坡度瞬间过零就立刻反转。

2.2 公式

引入一个"速度"变量 mm(初始为0),用指数移动平均的方式融合历史梯度:

mt=β⋅mt−1+(1−β)⋅gtm_t = \beta \cdot m_{t-1} + (1-\beta) \cdot g_t θt=θt−1−η⋅mt\theta_t = \theta_{t-1} - \eta \cdot m_t

其中 gt=∇θLg_t = \nabla_\theta L 是当前梯度,β\beta(常取0.9)控制"记多久之前的历史"——β\beta 越大,惯性越强,单次的梯度抖动越容易被历史平均掉。注意参数更新用的是 mtm_t,不是 gtg_t。

2.3 猫吃鱼例子:手算三步,看清"滞后性"

沿用"猫吃鱼"里第二层某个权重参数,取上面那组梯度序列的前三步:g1=+0.4, g2=0, g3=−0.4g_1=+0.4,\ g_2=0,\ g_3=-0.4。取 β=0.9\beta=0.9,m0=0m_0=0:

第1步:

m1=0.9×0+0.1×0.4=+0.04m_1 = 0.9\times0 + 0.1\times0.4 = +0.04

第2步:

m2=0.9×0.04+0.1×0=+0.036m_2 = 0.9\times0.04 + 0.1\times0 = +0.036

当前梯度已经是0了,但动量仍然是正的——历史趋势还在起作用,更新并没有立刻停下来。

第3步:

m3=0.9×0.036+0.1×(−0.4)=0.0324−0.04=−0.0076m_3 = 0.9\times0.036 + 0.1\times(-0.4) = 0.0324 - 0.04 = -0.0076

到这一步动量才翻负,更新方向才真正开始反转。

这就是Momentum最关键的两个特性:滞后性——不会跟着瞬时梯度立刻反向;记忆性——梯度归零的那一刻,更新仍然带着之前积累的趋势往前走。真正容易让朴素梯度下降剧烈震荡的场景,是损失函数地形像"峡谷"一样——某个方向陡、某个方向缓,每一步都冲过头再被拉回来;上面这组"梯度平滑过零"的例子,更适合用来直观感受动量的滞后性,而不是等价于峡谷震荡那种更剧烈的情形。

(一个记号上的提醒:公式是 θt=θt−1−ηmt\theta_t=\theta_{t-1}-\eta m_t,所以 mtm_t 为正时,参数是按公式方向持续更新,而不是简单地"mm正就往正方向走"——具体是加还是减,要看公式里那个负号。)

三、RMSProp:让每个参数有自己的步幅

3.1 直觉

Momentum解决的是"方向震荡",但没解决"不同参数该用不同学习率"这件事。RMSProp的思路是:给每个参数维护一个"历史梯度有多大"的统计量,梯度一直很大的参数,自动把步子迈小一点;梯度一直很小的参数,自动把步子迈大一点。

3.2 公式

维护梯度平方的指数移动平均 ss(初始为0):

st=β2⋅st−1+(1−β2)⋅gt2s_t = \beta_2 \cdot s_{t-1} + (1-\beta_2) \cdot g_t^2 θt=θt−1−ηst+ϵ⋅gt\theta_t = \theta_{t-1} - \frac{\eta}{\sqrt{s_t} + \epsilon} \cdot g_t

这里 ϵ\epsilon(通常 10−810^{-8})只是为了防止分母为0。关键在分母 st\sqrt{s_t}:如果某参数历史梯度一直很大,sts_t 就大,除下来实际步长就变小;反之步长变大。这就是"自适应学习率"的由来——学习率 η\eta 本身没变,但每个参数实际感受到的步长被 sts_t 重新缩放了。

比起Momentum,RMSProp直观得多:梯度大的参数除以一个大数,步子变小;梯度小的参数除以一个小数,步子变大。就是"除以一个跟历史梯度大小相关的数",没有滞后、没有过渡、没有方向翻转那些弯弯绕。

四、Adam:Momentum + RMSProp

4.1 完整公式

Adam(Adaptive Moment Estimation)同时维护两个统计量:一阶矩 mm(梯度的均值,即Momentum那一套)和二阶矩 vv(梯度平方的均值,即RMSProp那一套):

mt=β1⋅mt−1+(1−β1)⋅gtm_t = \beta_1 \cdot m_{t-1} + (1-\beta_1) \cdot g_t vt=β2⋅vt−1+(1−β2)⋅gt2v_t = \beta_2 \cdot v_{t-1} + (1-\beta_2) \cdot g_t^2

由于 m0=v0=0m_0 = v_0 = 0,训练刚开始的几步里 mt,vtm_t, v_t 会明显偏向0(尤其 β1,β2\beta_1,\beta_2 接近1时更明显)。Adam额外做了一步偏差修正(bias correction):

m^t=mt1−β1t,v^t=vt1−β2t\hat{m}_t = \frac{m_t}{1-\beta_1^t}, \qquad \hat{v}_t = \frac{v_t}{1-\beta_2^t}

最终更新:

θt=θt−1−ηv^t+ϵ⋅m^t\theta_t = \theta_{t-1} - \frac{\eta}{\sqrt{\hat{v}_t}+\epsilon} \cdot \hat{m}_t

默认超参数:β1=0.9\beta_1=0.9,β2=0.999\beta_2=0.999,ϵ=10−8\epsilon=10^{-8},η\eta 常取 10−310^{-3} 量级(对Transformer训练来说通常还要配合warmup,这是后面的话题)。

4.2 猫吃鱼例子:完整走一步

继续用第2节的梯度 g1=0.40g_1=0.40,取 β1=0.9,β2=0.999,η=0.1,ϵ=10−8\beta_1=0.9,\beta_2=0.999,\eta=0.1,\epsilon=10^{-8},从 t=1t=1 开始手算:

一阶矩:

m1=0.9×0+0.1×0.40=0.040m_1 = 0.9\times0 + 0.1\times0.40 = 0.040

二阶矩:

v1=0.999×0+0.001×0.402=0.001×0.16=0.00016v_1 = 0.999\times0 + 0.001\times0.40^2 = 0.001\times0.16 = 0.00016

偏差修正(注意 t=1t=1 时修正效果最明显):

m^1=0.0401−0.91=0.0400.1=0.40\hat{m}_1 = \frac{0.040}{1-0.9^1} = \frac{0.040}{0.1} = 0.40 v^1=0.000161−0.9991=0.000160.001=0.16\hat{v}_1 = \frac{0.00016}{1-0.999^1} = \frac{0.00016}{0.001} = 0.16

更新量:

Δθ=η⋅m^1v^1+ϵ=0.1×0.400.16+10−8=0.1×0.400.40≈0.1\Delta\theta = \eta \cdot \frac{\hat{m}_1}{\sqrt{\hat{v}_1}+\epsilon} = 0.1 \times \frac{0.40}{\sqrt{0.16}+10^{-8}} = 0.1 \times \frac{0.40}{0.40} \approx 0.1

有意思的地方来了:偏差修正之后,m^1\hat{m}_1 和 v^1\hat{v}_1 其实分别还原成了"未修正前的原始梯度尺度"(m^1=g1=0.40\hat{m}_1=g_1=0.40,v^1=∣g1∣=0.40\sqrt{\hat{v}_1}=|g_1|=0.40),两者相除正好约等于1,所以第一步的实际更新量约等于学习率本身。这不是巧合——偏差修正的设计目标之一就是让训练刚开始时的更新量不被 m0=v0=0m_0=v_0=0 的初始偏差压得过小。

五、直观理解:Adam在做什么

把两部分放一起看:

  • 分子 m^t\hat{m}_t(方向):告诉你"往哪走"——用历史梯度的平滑平均代替当前这一次可能带噪声的梯度,方向更稳。
  • 分母 v^t\sqrt{\hat{v}_t}(步幅):告诉你"走多远"——历史梯度波动大的参数除以一个更大的数,步子自动收窄;波动小的参数步子自动放开。

可以理解成:Momentum负责"不要被一次性的噪声带偏方向",RMSProp负责"给每个参数配一双合脚的鞋"。Adam把这两件事一起做了,这也是为什么它比朴素SGD在实践中通常更快收敛、对学习率的选择也更不敏感。

六、PyTorch实现

6.1 手写一步Adam更新(照着公式写,不调库)

import torch

# 沿用猫吃鱼例子里的某个权重参数
w = torch.tensor(0.50, requires_grad=True)

# Adam的状态:一阶矩、二阶矩、时间步
m = torch.zeros_like(w)
v = torch.zeros_like(w)
t = 0

beta1, beta2, eps, lr = 0.9, 0.999, 1e-8, 0.1

def adam_step(w, grad, m, v, t):
    t += 1
    m = beta1 * m + (1 - beta1) * grad
    v = beta2 * v + (1 - beta2) * grad ** 2
    m_hat = m / (1 - beta1 ** t)
    v_hat = v / (1 - beta2 ** t)
    w = w - lr * m_hat / (torch.sqrt(v_hat) + eps)
    return w, m, v, t

# 假设这是第1步反向传播算出的梯度(对应上面手算的例子)
grad = torch.tensor(0.40)
w, m, v, t = adam_step(w, grad, m, v, t)
print(w.item())  # 约等于 0.40,验证了上面的手算结果

6.2 用torch.optim.Adam训练"猫吃鱼"GPT

真实训练里没人会手写Adam,直接用PyTorch内置的优化器就行,接口和EP3里用的SGD完全一致,只需要换一行:

import torch.optim as optim

model = GPT(...)  # 推理篇 EP12 里搭好的完整 GPT 模型
optimizer = optim.Adam(model.parameters(), lr=1e-3, betas=(0.9, 0.999), eps=1e-8)

for step in range(num_steps):
    optimizer.zero_grad()          # 清空上一轮的梯度(否则会累加)
    logits = model(input_ids)      # 前向传播
    loss = loss_fn(logits, targets)
    loss.backward()                # 反向传播,算出所有参数的梯度
    optimizer.step()               # 这一步内部就是我们上面手写的Adam更新逻辑

optimizer.step() 这一行背后,PyTorch为模型里每一个参数都各自维护了一份 mm、vv、tt,逐参数独立地做我们上面手算的那套流程——这也是"自适应"的真正含义:不是所有参数共用一个学习率,而是每个参数都有自己的"步幅历史"。

七、小结

记住的东西 解决的问题
Momentum 梯度的历史均值 mm 方向震荡
RMSProp 梯度平方的历史均值 vv 学习率一刀切
Adam mm 和 vv 都记,外加偏差修正 两个问题一起解决

系列回顾

  • 训练篇 EP01:最小二乘法与损失函数
  • 训练篇 EP02:反向传播到底在传什么
  • 训练篇 EP03:梯度下降
  • 训练篇 EP04(本期):Adam优化器

预告

优化器选好了,更新规则也有了,但真实训练里学习率往往不是从头到尾固定不变的——训练初期太大容易震荡发散,训练后期太大又会在最优点附近来回跳。下一期训练篇 EP05,我们会聊聊学习率调度(Learning Rate Scheduling),把warmup、衰减这些训练里常见但容易被忽略的细节讲清楚。


本文为"从零手写大模型"系列训练篇第4期,系列完整代码见项目仓库。

← 返回训练篇目录