第 17 课:AdamW 优化器 —— 给梯度下降装上「惯性」和「自适应步长」
第 17 课:AdamW 优化器 —— 给梯度下降装上"惯性"和"自适应步长"
代码位置:src/optim.rs(
SGD/AdamW) 代码位置:src/train.rs(train_gpt中 AdamW 的用法) 演示入口:src/main.rs(演示 3:训练小 GPT)
1. 本课要搞懂的问题
- 朴素 SGD 有什么致命弱点?为什么深度模型上它不够用?
- Adam 的一阶动量 m、二阶动量 v 分别"记忆"了什么?对应公式怎么写?
- 为什么需要偏差修正(除以 1-β^t)?前几步不修会怎样?
- 权重衰减"解耦"(decoupled weight decay)和传统的 L2 正则化有什么区别?
AdamW::step的每一行代码分别对应公式里的哪一项?
2. 先复习:SGD 的局限
第 6 课实现的 SGD 是整个故事的原点:
/// 更新一步:θ = θ - lr * g
pub fn step(&self) {
for p in &self.params {
let g = p.grad();
let d = p.data();
let updated: Vec<f32> = d.iter().zip(&g).map(|(v, g)| v - self.lr * g).collect();
p.set_data(updated);
}
}
公式只有一个:
θ ← θ - lr · g
它有两个致命弱点:
| 弱点 | 表现 | 直观理解 |
|---|---|---|
| 梯度震荡 | 参数在最优解附近来回横跳,收敛慢 | 下山时每一步都只看"当前脚下"的坡度,方向忽左忽右,没有"惯性" |
| 步长不会自适应 | 所有参数共用同一个 lr | 有的维度坡缓(需要大步长),有的维度坡陡/噪声大(需要小步长),SGD 一刀切 |
现实世界的梯度方向总是带噪声的(batch 采样引入的随机性、损失面本身的崎岖),只信"当前这一下"的 SGD 在深度网络上要么太慢、要么震荡发散。
3. Adam 的核心思想:一阶动量 + 二阶动量
Adam(Adaptive Moment Estimation)给每个参数维护两个额外状态,公式如下:
一阶动量(梯度的指数移动平均,记住"方向"):
m_t = β1 · m_{t-1} + (1 - β1) · g_t
二阶动量(梯度平方的指数移动平均,感知"陡峭程度"):
v_t = β2 · v_{t-1} + (1 - β2) · g_t²
参数更新:
θ_t = θ_{t-1} - lr · m̂_t / (√v̂_t + ε)
直觉对应:
| 状态 | 记忆的内容 | 类比 | 作用 |
|---|---|---|---|
一阶动量 m |
梯度的平均值(带方向) | 小球下坡的惯性:过去几步都往东偏,这次也往东多走点 | 平滑梯度、抵消震荡 |
二阶动量 v |
梯度平方的平均值(无方向,恒正) | 对坡度的感知:某个维度长期坡陡,说明这里"水深",步子要小 | 每个参数独立缩放步长 |
注意 m、v 都是逐参数、逐元素维护的(AdamW 里 m: Vec<Vec<f32>>、v: Vec<Vec<f32>>,和每个参数张量形状一一对应),所以"自适应步长"是精细到每个标量权重的。
为什么用指数移动平均而不是简单平均
因为只需要记住 m_{t-1} 和 v_{t-1} 两个状态就能增量更新,不需要存全部历史梯度;β1=0.9 意味着过去约 10 步的梯度主导,β2=0.999 意味着过去约 1000 步的梯度平方主导——越靠前的历史衰减得越厉害。
4. 偏差修正:除以 (1 - β^t)
问题:m_0 = 0、v_0 = 0,训练第一步 m_1 = (1-β1)·g_1,只有真实梯度的 10%(β1=0.9 时)。训练初期 m、v 被严重"低估",直接使用会导致起步步长偏小。
修正办法:把 m_t、v_t 除以各自的累积衰减系数:
m̂_t = m_t / (1 - β1^t)
v̂_t = v_t / (1 - β2^t)
| t | 1 - β1^t(β1=0.9) | 1 - β2^t(β2=0.999) | 效果 |
|---|---|---|---|
| 1 | 0.1 | 0.001 | 修正最狠:m̂ = m/0.1 = 10×m,v̂ = v/0.001 = 1000×v |
| 10 | 1 - 0.9¹⁰ ≈ 0.651 | 1 - 0.999¹⁰ ≈ 0.00995 | 仍在修正 |
| 100 | 1 - 0.9¹⁰⁰ ≈ 1.0 | 1 - 0.999¹⁰⁰ ≈ 0.095 | m 基本不用修,v 还要修 |
| 1000 | ≈ 1.0 | ≈ 0.632 | v 接近不需要修 |
关键观察:
t越大,β^t越接近 0,修正系数越接近 1——偏差修正只在训练初期起作用。而且m的 β1 小、修正结束得快,v的 β2 接近 1、修正持续得更久。这正是代码里两个系数分开算的原因。
5. 权重衰减解耦(decoupled weight decay)
正则化思想:每步更新时,额外把参数往 0 拉一点点,防止权重过大、过拟合。
传统做法(L2 正则化):把 λ·θ 加进损失再求导,相当于梯度变成 g + λ·θ,然后被 Adam 的"自适应步长"一通缩放——λ·θ 这个衰减项也被 1/√v̂ 缩放了,衰减强度随梯度历史变化,不受控制。
AdamW 的解耦做法:权重衰减独立于梯度,作为一个单独的减法项:
θ ← θ - lr · m̂/(√v̂ + ε) - lr · wd · θ
└──── Adam 步长 ────┘ └─ 解耦的权重衰减 ─┘
| 传统 L2(Adam+L2) | AdamW(decoupled) | |
|---|---|---|
| 衰减项怎么来 | 进损失函数求导,混进梯度 g+λθ |
不进梯度,更新时单独减 lr·wd·θ |
| 是否被自适应步长缩放 | 是(被 1/√v̂ 缩放,强度不稳定) |
否(恒定 lr·wd,与梯度历史无关) |
| 实际效果 | 衰减幅度时大时小,难调 | 衰减可预期、好调超参 |
在现代 LLM 训练里(GPT 系列等)几乎都用 AdamW,就是因为这个"可预期"的衰减。
6. AdamW::step 逐行讲解
完整代码(src/optim.rs):
pub fn step(&mut self) {
self.t += 1;
// 偏差修正系数(训练初期 t 小,修正大)
let bc1 = 1.0 - self.beta1.powi(self.t as i32);
let bc2 = 1.0 - self.beta2.powi(self.t as i32);
for (i, p) in self.params.iter().enumerate() {
let g = p.grad();
let d = p.data();
let mut updated = vec![0.0f32; d.len()];
for j in 0..d.len() {
let gv = g[j];
// 1. 更新动量
self.m[i][j] = self.beta1 * self.m[i][j] + (1.0 - self.beta1) * gv;
self.v[i][j] = self.beta2 * self.v[i][j] + (1.0 - self.beta2) * gv * gv;
// 2. 偏差修正
let m_hat = self.m[i][j] / bc1;
let v_hat = self.v[i][j] / bc2;
// 3. 更新:θ -= lr * m_hat/(√v_hat + eps) + lr * wd * θ(权重衰减解耦)
let step = self.lr * m_hat / (v_hat.sqrt() + self.eps);
let decay = self.lr * self.weight_decay * d[j];
updated[j] = d[j] - step - decay;
}
p.set_data(updated);
}
}
逐行对应公式:
| 代码 | 对应公式 | 说明 |
|---|---|---|
self.t += 1; |
—— | 步数计数,偏差修正要用到 t |
bc1 = 1.0 - self.beta1.powi(t) |
1 - β1^t |
一阶动量修正系数(t 从 1 开始) |
bc2 = 1.0 - self.beta2.powi(t) |
1 - β2^t |
二阶动量修正系数 |
self.m[i][j] = self.beta1 * self.m[i][j] + (1.0 - self.beta1) * gv; |
m_t = β1·m_{t-1} + (1-β1)·g_t |
一阶动量:新旧梯度按 9:1 加权 |
self.v[i][j] = self.beta2 * self.v[i][j] + (1.0 - self.beta2) * gv * gv; |
v_t = β2·v_{t-1} + (1-β2)·g_t² |
二阶动量:梯度平方,恒正 |
m_hat = self.m[i][j] / bc1; |
m̂_t = m_t/(1-β1^t) |
修正初期被低估的一阶动量 |
v_hat = self.v[i][j] / bc2; |
v̂_t = v_t/(1-β2^t) |
修正初期被低估的二阶动量 |
step = self.lr * m_hat / (v_hat.sqrt() + self.eps); |
lr·m̂/(√v̂+ε) |
Adam 更新步长:方向 m̂,大小被 √v̂ 自适应缩放 |
decay = self.lr * self.weight_decay * d[j]; |
lr·wd·θ |
解耦的权重衰减,不经过 √v̂ 缩放 |
updated[j] = d[j] - step - decay; |
θ ← θ - step - decay |
参数更新 = 旧值 - Adam 步长 - 衰减 |
p.set_data(updated); |
—— | 写回参数张量 |
四个值得停下来想一想的细节:
- 为什么
v_hat.sqrt()要加eps:防止v_hat ≈ 0(训练初期)时除零,数值稳定性用。eps = 1e-8是 Adam 论文的标准值。 lr是pub字段:src/train.rs里每步都改它——
调度器只负责改opt.lr = scheduler.lr(); // 第 20 课:warmup + cosine 调度 opt.step();lr,AdamW 的m、v状态跨步累积、完全不受影响。- 权重衰减项用的是
d[j](更新前的旧参数):这就是"解耦"——衰减直接作用在参数本身,而不是作用在梯度上。 m、v与参数逐元素对齐:AdamW::new里params.iter().map(|p| vec![0.0f32; p.numel()]),每个参数张量配一个同长度的一阶/二阶动量数组,更新时按j同步遍历。
7. AdamW 的超参数
AdamW::new 里的默认值就是论文标准:
AdamW {
lr,
beta1: 0.9,
beta2: 0.999,
eps: 1e-8,
weight_decay,
t: 0,
params,
m,
v,
}
| 超参数 | 值 | 含义 | 调参经验 |
|---|---|---|---|
beta1 |
0.9 | 一阶动量衰减系数 | 一般不动,范围 0.8~0.95 |
beta2 |
0.999 | 二阶动量衰减系数 | 一般不动;训练不稳时有人降到 0.95 |
eps |
1e-8 | 除零保护 | 一般不动 |
weight_decay |
0.01(本项目) | 权重衰减强度 | 常用 0.01~0.1;越大正则化越强 |
lr |
调度器控制 | 基础步长 | 配合 warmup/cosine(第 20 课) |
本项目在 train_gpt 里这样接入:
let mut opt = AdamW::new(max_lr, params.clone(), 0.01); // wd = 0.01
...
clip_grad_norm(¶ms, 1.0); // 先裁剪梯度(防止个别大梯度冲坏 m、v 的估计)
opt.lr = scheduler.lr(); // 再设置当前学习率
opt.step(); // 最后更新
opt.zero_grad();
顺序值得注意:梯度裁剪在
AdamW::step之前。因为 Adam 的 m、v 是梯度历史的长期记忆,如果某步梯度爆炸没被裁掉,它会污染 m/v 很久。先裁剪、再进优化器,是 LLM 训练的标准顺序。
8. 动手练习
- 对比 SGD vs AdamW:临时把
train_gpt里的AdamW::new(...)换成SGD::new(0.01, params.clone())(去掉opt.lr = scheduler.lr()那行,因为SGD::lr是私有字段),跑 600 步看 loss——体会 AdamW 在小语料上的收敛速度优势。 - 手推第一步:假设某参数
θ=1.0、g=0.5、lr=0.001、wd=0.01,手算t=1时m_hat、v_hat、step、decay,再在AdamW::step里加一行println!验证。 - 看偏差修正的效果:把
bc1、bc2改成恒为 1.0(不修正),训练对比 loss 曲线——训练初期应该明显变慢。 - 调权重衰减:把
wd从 0.01 改成 0.1 和 0.0,观察训练后 loss 与生成文本的差异,体会正则化的作用。 - 思考:为什么
v存的是g²而不是|g|?如果某个维度的梯度长期是 +0.1/-0.1 交替(震荡),m和v分别是什么表现?(提示:m 会被抵消趋近 0,v 会累积为正——这就是 Adam 抑制震荡的机制。)
9. 本课总结
-
SGD 只有
θ -= lr·g,既没有惯性(震荡)也不能自适应步长(一刀切) -
Adam 用一阶动量
m = β1·m + (1-β1)·g平滑方向,用二阶动量v = β2·v + (1-β2)·g²感知陡峭度,更新为θ -= lr·m̂/(√v̂+ε) -
偏差修正除以
1-β^t:只在训练初期起作用,补偿 m/v 从 0 起步的低估 -
AdamW 把权重衰减解耦成独立项
θ -= lr·wd·θ,不受自适应步长干扰,衰减强度可预期 -
AdamW::step的每行代码都能在公式里找到对应项:t→ 修正系数 → 逐元素动量 → 修正 → 步长/衰减 → 写回 -
下一课:推理加速神器 KV Cache——为什么生成时要缓存 K/V,怎么做到"每步只算 1 个 token"。