从零手写大模型 · 训练篇 04 - Adam Optimizer(Adam 优化器)
训练篇 · 查看系列上期回顾
在训练篇 EP03里,我们把梯度下降(Gradient Descent)的更新规则从"感觉上应该往梯度反方向走"落地成了一行公式:
然后沿用"猫吃鱼"这个2层小例子,手动做了一步更新(学习率 ),loss从约 0.1914 降到了约 0.0157。我们也顺手过了一遍batch / mini-batch / SGD的区别。
但如果你把这个更新规则原封不动地套到真实训练里,很快会撞见两个很现实的问题。这就是今天的起点。
一、朴素梯度下降的两个痛点
痛点一:方向会抖。
想象损失函数的地形不是一个规规矩矩的碗,而是一条窄窄的峡谷——某个方向陡,某个方向缓。朴素梯度下降在陡的方向上会来回震荡(每步都冲过头,再被拉回来),在缓的方向上却挪得特别慢。整个下降路径像蛇形一样左右摆动,而不是径直冲向谷底。
痛点二:所有参数共用一个学习率。
在一个几十万甚至上亿参数的模型里,不同参数的梯度量级可能差好几个数量级。有的参数梯度一直很大,学习率稍微设大一点就会震荡发散;有的参数梯度一直很小,同样的学习率下几乎不动。用同一个 去伺候所有参数,注定是"一刀切"式的妥协。
这两个痛点分别对应两类改进思路:动量(Momentum)解决"抖",自适应学习率(如RMSProp)解决"一刀切"。Adam把这两者结合在了一起,这也是它至今仍是深度学习默认优化器的原因。
二、Momentum:给梯度下降加上惯性
2.1 直觉
朴素梯度下降每一步只看"当前这一刻"的梯度,完全没有记忆。假设某个参数在训练中,梯度一开始是正的,然后慢慢减小、越过零点、变成负的:
如果没有动量,梯度一变负,参数立刻就要往反方向走——这就容易在最低点附近来回过头、来回震荡。Momentum的想法很简单:让更新量带上一点"惯性",不要只听当前这一步梯度的,也要参考之前一直积累的趋势。
物理类比很贴切:一个小球从山坡滚下来,如果只受"当前位置的坡度"支配,遇到小坑洼就会卡住来回晃;但如果小球有质量、有惯性,它会在整体下坡方向上越滚越快,局部的小震荡会被惯性"平均"掉,方向也不会因为坡度瞬间过零就立刻反转。
2.2 公式
引入一个"速度"变量 (初始为0),用指数移动平均的方式融合历史梯度:
其中 是当前梯度,(常取0.9)控制"记多久之前的历史"—— 越大,惯性越强,单次的梯度抖动越容易被历史平均掉。注意参数更新用的是 ,不是 。
2.3 猫吃鱼例子:手算三步,看清"滞后性"
沿用"猫吃鱼"里第二层某个权重参数,取上面那组梯度序列的前三步:。取 ,:
第1步:
第2步:
当前梯度已经是0了,但动量仍然是正的——历史趋势还在起作用,更新并没有立刻停下来。
第3步:
到这一步动量才翻负,更新方向才真正开始反转。
这就是Momentum最关键的两个特性:滞后性——不会跟着瞬时梯度立刻反向;记忆性——梯度归零的那一刻,更新仍然带着之前积累的趋势往前走。真正容易让朴素梯度下降剧烈震荡的场景,是损失函数地形像"峡谷"一样——某个方向陡、某个方向缓,每一步都冲过头再被拉回来;上面这组"梯度平滑过零"的例子,更适合用来直观感受动量的滞后性,而不是等价于峡谷震荡那种更剧烈的情形。
(一个记号上的提醒:公式是 ,所以 为正时,参数是按公式方向持续更新,而不是简单地"正就往正方向走"——具体是加还是减,要看公式里那个负号。)
三、RMSProp:让每个参数有自己的步幅
3.1 直觉
Momentum解决的是"方向震荡",但没解决"不同参数该用不同学习率"这件事。RMSProp的思路是:给每个参数维护一个"历史梯度有多大"的统计量,梯度一直很大的参数,自动把步子迈小一点;梯度一直很小的参数,自动把步子迈大一点。
3.2 公式
维护梯度平方的指数移动平均 (初始为0):
这里 (通常 )只是为了防止分母为0。关键在分母 :如果某参数历史梯度一直很大, 就大,除下来实际步长就变小;反之步长变大。这就是"自适应学习率"的由来——学习率 本身没变,但每个参数实际感受到的步长被 重新缩放了。
比起Momentum,RMSProp直观得多:梯度大的参数除以一个大数,步子变小;梯度小的参数除以一个小数,步子变大。就是"除以一个跟历史梯度大小相关的数",没有滞后、没有过渡、没有方向翻转那些弯弯绕。
四、Adam:Momentum + RMSProp
4.1 完整公式
Adam(Adaptive Moment Estimation)同时维护两个统计量:一阶矩 (梯度的均值,即Momentum那一套)和二阶矩 (梯度平方的均值,即RMSProp那一套):
由于 ,训练刚开始的几步里 会明显偏向0(尤其 接近1时更明显)。Adam额外做了一步偏差修正(bias correction):
最终更新:
默认超参数:,,, 常取 量级(对Transformer训练来说通常还要配合warmup,这是后面的话题)。
4.2 猫吃鱼例子:完整走一步
继续用第2节的梯度 ,取 ,从 开始手算:
一阶矩:
二阶矩:
偏差修正(注意 时修正效果最明显):
更新量:
有意思的地方来了:偏差修正之后, 和 其实分别还原成了"未修正前的原始梯度尺度"(,),两者相除正好约等于1,所以第一步的实际更新量约等于学习率本身。这不是巧合——偏差修正的设计目标之一就是让训练刚开始时的更新量不被 的初始偏差压得过小。
五、直观理解:Adam在做什么
把两部分放一起看:
- 分子 (方向):告诉你"往哪走"——用历史梯度的平滑平均代替当前这一次可能带噪声的梯度,方向更稳。
- 分母 (步幅):告诉你"走多远"——历史梯度波动大的参数除以一个更大的数,步子自动收窄;波动小的参数步子自动放开。
可以理解成: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为模型里每一个参数都各自维护了一份 、、,逐参数独立地做我们上面手算的那套流程——这也是"自适应"的真正含义:不是所有参数共用一个学习率,而是每个参数都有自己的"步幅历史"。
七、小结
| 记住的东西 | 解决的问题 | |
|---|---|---|
| Momentum | 梯度的历史均值 | 方向震荡 |
| RMSProp | 梯度平方的历史均值 | 学习率一刀切 |
| Adam | 和 都记,外加偏差修正 | 两个问题一起解决 |
系列回顾
- 训练篇 EP01:最小二乘法与损失函数
- 训练篇 EP02:反向传播到底在传什么
- 训练篇 EP03:梯度下降
- 训练篇 EP04(本期):Adam优化器
预告
优化器选好了,更新规则也有了,但真实训练里学习率往往不是从头到尾固定不变的——训练初期太大容易震荡发散,训练后期太大又会在最优点附近来回跳。下一期训练篇 EP05,我们会聊聊学习率调度(Learning Rate Scheduling),把warmup、衰减这些训练里常见但容易被忽略的细节讲清楚。
本文为"从零手写大模型"系列训练篇第4期,系列完整代码见项目仓库。