第 17 课:AdamW 优化器 —— 给梯度下降装上「惯性」和「自适应步长」

2026-09-18 干徒
RustLLM优化器

第 17 课:AdamW 优化器 —— 给梯度下降装上"惯性"和"自适应步长"

代码位置:src/optim.rsSGD / AdamW) 代码位置:src/train.rstrain_gpt 中 AdamW 的用法) 演示入口:src/main.rs(演示 3:训练小 GPT)

1. 本课要搞懂的问题

  1. 朴素 SGD 有什么致命弱点?为什么深度模型上它不够用?
  2. Adam 的一阶动量 m、二阶动量 v 分别"记忆"了什么?对应公式怎么写?
  3. 为什么需要偏差修正(除以 1-β^t)?前几步不修会怎样?
  4. 权重衰减"解耦"(decoupled weight decay)和传统的 L2 正则化有什么区别?
  5. 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 梯度平方的平均值(无方向,恒正) 对坡度的感知:某个维度长期坡陡,说明这里"水深",步子要小 每个参数独立缩放步长

注意 mv 都是逐参数、逐元素维护的(AdamWm: 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 = 0v_0 = 0,训练第一步 m_1 = (1-β1)·g_1,只有真实梯度的 10%(β1=0.9 时)。训练初期 mv 被严重"低估",直接使用会导致起步步长偏小。

修正办法:把 m_tv_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×mv̂ = 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); —— 写回参数张量

四个值得停下来想一想的细节:

  1. 为什么 v_hat.sqrt() 要加 eps:防止 v_hat ≈ 0(训练初期)时除零,数值稳定性用。eps = 1e-8 是 Adam 论文的标准值。
  2. lrpub 字段src/train.rs 里每步都改它——
    opt.lr = scheduler.lr();   // 第 20 课:warmup + cosine 调度
    opt.step();
    
    调度器只负责改 lr,AdamW 的 mv 状态跨步累积、完全不受影响。
  3. 权重衰减项用的是 d[j](更新前的旧参数):这就是"解耦"——衰减直接作用在参数本身,而不是作用在梯度上。
  4. mv 与参数逐元素对齐AdamW::newparams.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(&params, 1.0);   // 先裁剪梯度(防止个别大梯度冲坏 m、v 的估计)
opt.lr = scheduler.lr();        // 再设置当前学习率
opt.step();                     // 最后更新
opt.zero_grad();

顺序值得注意:梯度裁剪在 AdamW::step 之前。因为 Adam 的 m、v 是梯度历史的长期记忆,如果某步梯度爆炸没被裁掉,它会污染 m/v 很久。先裁剪、再进优化器,是 LLM 训练的标准顺序。

8. 动手练习

  1. 对比 SGD vs AdamW:临时把 train_gpt 里的 AdamW::new(...) 换成 SGD::new(0.01, params.clone())(去掉 opt.lr = scheduler.lr() 那行,因为 SGD::lr 是私有字段),跑 600 步看 loss——体会 AdamW 在小语料上的收敛速度优势。
  2. 手推第一步:假设某参数 θ=1.0g=0.5lr=0.001wd=0.01,手算 t=1m_hatv_hatstepdecay,再在 AdamW::step 里加一行 println! 验证。
  3. 看偏差修正的效果:把 bc1bc2 改成恒为 1.0(不修正),训练对比 loss 曲线——训练初期应该明显变慢。
  4. 调权重衰减:把 wd 从 0.01 改成 0.1 和 0.0,观察训练后 loss 与生成文本的差异,体会正则化的作用。
  5. 思考:为什么 v 存的是 而不是 |g|?如果某个维度的梯度长期是 +0.1/-0.1 交替(震荡),mv 分别是什么表现?(提示: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"。

Rust 大语言模型 学习指南共 22 章
1从零用 Rust 实现大语言模型 —— 学习计划2第 1 课:张量 Tensor —— 一切的基础3第 2 课:自动微分 Autograd —— 让模型学会「自我修正」4第 3 课:张量运算扩展 —— 广播、归约、softmax、批量矩阵乘法5第 4 课:模块化重构 —— 项目结构分层与 Rc<RefCell> 架构6第 5 课:线性层与激活函数 —— 神经网络的「积木」7第 6 课:损失函数与优化器 —— 让模型知道「错在哪、怎么改」8第 7 课:第一个 MLP —— 教会神经网络算 XOR9第 8 课:BPE 分词器 —— 让模型「读懂」文字10第 9 课:注意力机制 —— 让 token 互相「看」11第 10 课:多头注意力 —— 让模型「多角度」看世界12第 11 课:位置编码与归一化 —— 让序列带上「位置感」13第 12 课:完整 GPT 模型 —— 把积木拼成能预测下一个词的模型14第 13 课:训练循环 —— 让模型真正开始学习15第 14 课:数据加载 —— 文本如何变成训练样本16第 15 课:推理与采样 —— 让模型「创造性」地生成17第 16 课:训练小 GPT —— 看 loss 从 1.46 一路降到 0.1618第 17 课:AdamW 优化器 —— 给梯度下降装上「惯性」和「自适应步长」本篇19第 18 课:KV Cache —— 让逐 token 生成不再重复计算20第 19 课:RoPE 旋转位置编码 —— 把「相对位置」揉进注意力21第 20 课:学习率调度与收尾 —— warmup、cosine decay 与全项目总结22第 21 课:GPU 加速训练与推理