学习率调度深度实战:从 Warmup、Cosine 退火到 WSD 与线性缩放律的大批量预训练工程

系列定位:本文与站内《RMSNorm 深度实战》《相对位置编码深度实战》《交叉注意力深度实战》共同构成「Transformer 内部机制与训练工程深度解」系列。在 RMSNorm 解决了归一化的数值稳定性、相对位置编码解决了长度外推之后,本文聚焦 optimizer 与学习率调度这一训练动力学的灵魂——它决定了模型能否收敛到泛化良好的平坦极小值,而非尖锐的过拟合洼地。

1. 为什么学习率调度是训练的第一性原理

学习率 η 是神经网络训练中唯一一个 没有梯度、却主宰整个优化轨迹 的超参数。对带动量(momentum)的 SGD 更新:


v_{t+1} = β · v_t + (1 - β) · ∇L(θ_t)
θ_{t+1} = θ_t - η · v_{t+1}

η 同时控制「步长」与「势能盆地里的逃逸能力」。固定 η 在训练中几乎总是次优:

  • 训练初期:参数在随机初始化附近,loss 曲面高度非凸、梯度方向噪声极大。若 η 过大,参数会剧烈震荡甚至发散;若过小,则浪费前几千步的「探索」窗口。
  • 训练后期:loss 已接近一个宽而浅的 basin,需要 η 缩小到能将轨迹「锁」进极小值附近,否则会在极小值周边持续抖动、无法进一步下降。
  • 泛化视角:大量实证(如《SGD 的隐式偏置》《sharp vs flat minima》)表明,合适的退火 倾向于收敛到 flat minima,其对数据扰动的鲁棒性显著优于 sharp minima——而 flat minima 的抵达高度依赖 η 的收敛路径。

因此「调度」的本质不是玄学,而是 在训练轨迹的不同阶段,主动改变优化器的有效步长,以匹配不断变化的 loss 几何。

2. 恒定与阶梯衰减:朴素基线

最朴素的做法是恒定学习率,或「阶梯衰减」(Step Decay):


# 阶梯衰减:每过 decay_step 个 epoch,学习率乘以 gamma
def step_decay(epoch, base_lr=1e-3, decay_step=30, gamma=0.1):
    return base_lr * (gamma ** (epoch // decay_step))

阶梯衰减在 ResNet/CNN 时代是事实标准(如原始 ResNet 在 30/60/90 epoch 处 ×0.1),但存在两个硬伤:

  1. 突变不连续:LR 在断点处瞬间跳变,优化轨迹会在断点附近产生可见的 loss 台阶(见图 2 的台阶状曲线)。
  2. 与阶段语义脱节:衰减时机靠人工拍脑袋,无法自适应 loss 收敛速度。

现代大模型训练已几乎全面放弃纯阶梯衰减,转向「平滑退火 + 预热」组合。

3. 预热(Warmup):为什么大模型训练离不开它

Warmup 指在训练最开始的 T_warm 步内,将学习率从 0(或极小值)线性/常数/余弦地爬升到目标峰值 η_max,再进入退火阶段。

3.1 为什么需要 Warmup

学术界给出过多个互补的解释,工程上最关键的三个:

成因 机制 不做 Warmup 的后果
梯度噪声方差随批量增大而放大 大批量下每步梯度估计方差小、方向更「硬」,起点高 η 易把参数甩出 basin 训练 loss 爆炸(NaN/Inf)
Adam 二阶矩 v_t 冷启动偏差 初期 v_t 远小于真实二阶矩,biased m̂_t/√v̂_t 被放大,等效步长失控 早期参数剧烈跳动
归一化层统计量未稳定 BatchNorm/LayerNorm 的 running statistics 在首几步不可靠 激活值分布漂移,梯度方向噪声大

3.2 三种 Warmup 形态


def linear_warmup(step, warmup_steps, peak_lr):
    return peak_lr * min(step, warmup_steps) / max(1, warmup_steps)

def constant_warmup(step, warmup_steps, peak_lr):
    # UL2 / PaLM 风格:前段用恒定小 LR,再跳到峰值
    return peak_lr * (0.1 if step < warmup_steps else 1.0)

def cosine_warmup(step, warmup_steps, peak_lr):
    if step >= warmup_steps:
        return peak_lr
    return peak_lr * (1 - math.cos(math.pi * step / warmup_steps)) / 2

经验法则:warmup_steps 通常设为总训练步数的 0.5%~3%;LLaMA-7B 用 2000 步 warmup,GPT-3 用 ~375M token 对应步数。

4. 余弦退火与 Cosine-with-Warmup

Cosine Annealing 让学习率沿半余弦曲线从 η_max 平滑下降到 η_min:


η_t = η_min + 0.5 · (η_max - η_min) · (1 + cos(π · t / T_total))

配合 warmup 的组合(HuggingFace get_cosine_schedule_with_warmup)已成为 Encoder 微调与中小模型预训练的事实标准:


import math
def cosine_with_warmup(step, warmup_steps, total_steps, peak_lr, min_lr=1e-5):
    if step < warmup_steps:
        return peak_lr * step / max(1, warmup_steps)
    progress = (step - warmup_steps) / max(1, total_steps - warmup_steps)
    progress = min(progress, 1.0)
    return min_lr + 0.5 * (peak_lr - min_lr) * (1 + math.cos(math.pi * progress))

为何余弦优于线性衰减? 余弦在退火中段下降快、末段趋平,给「中期快速穿越 basin、末期精细收敛」提供了更自然的节奏。SGDR(余弦带重启)则周期性把 LR 拉回峰值,用于逃离尖锐极小值、做快照集成。

5. Warmup-Stable-Decay(WSD):现代 LLM 预训练范式

到了百亿/千亿参数预训练, cosine 全程衰减暴露一个工程痛点:衰减过早开始会浪费大量「 Plateau 稳定训练」的算力。LLaMA-2、DeepSeek、Qwen 等采用的 WSD(Warmup-Stable-Decay) 把训练切成三段:

  1. Warmup:爬升到 η_max(同前)。
  2. Stable(恒定段):在 η_max 恒定训练绝大部分步数——这是算力利用率最高的阶段。
  3. Decay:最后 ~5%~10% 步数内,将 η_max 快速退火到 η_min(常用 cosine 或线性 decay)。

def wsd_schedule(step, warmup_steps, decay_start, total_steps, peak_lr, min_lr=1e-6):
    if step < warmup_steps:
        return peak_lr * step / max(1, warmup_steps)
    if step < decay_start:
        return peak_lr                      # stable 段:恒定
    # decay 段:在剩余步数内 cosine 退火
    decay_progress = (step - decay_start) / max(1, total_steps - decay_start)
    decay_progress = min(decay_progress, 1.0)
    return min_lr + 0.5 * (peak_lr - min_lr) * (1 + math.cos(math.pi * decay_progress))

WSD 的工程收益:

  • 可中途续训(continue training):stable 段恒定,意味着在 decay 之前任何时候停下、换数据继续训,都不会损失「退火进度」——这支撑了「训练到一半发现数据不够、补数据续训」的真实工作流。
  • decay 比例可调:实验表明最后 10% 步数的 decay 贡献了可观的 loss 下降(所谓 "late decay" 现象),把 decay 压缩到尾部反而更省算力。

6. 线性缩放律与 Warmup 修正:大批量训练的铁律

当用数据并行把 batch size 从 B 扩大到 kB 时,线性缩放律(Linear Scaling Rule) 指出:为保持「每样本等效更新」不变,峰值学习率应同步放大 k 倍:


η(kB) ≈ k · η(B)

直观理解:大批量每步见了 k 倍样本,等效于小批量走了 k 步,所以单步 LR 也要 k 倍才能匹配轨迹。

关键修正——Accumulated Warmup:线性缩放只在 warmup 结束之后成立。warmup 期间若直接放大 k 倍峰值,起点仍是 0,但 到达峰值所需的绝对步数不变 会导致大批量在 warmup 内「等效探索步数」不足。因此实践中常让 warmup_steps 也随 k 线性放大:


warmup_steps(kB) ≈ k · warmup_steps(B)

否则会出现「大批量 + 高 LR + 短 warmup」三连击,直接训练发散——这是分布式训练最常见的隐性 bug 之一。

7. 与梯度累积、混合精度、梯度裁剪的交互陷阱

学习率调度从来不是孤立的,它和训练栈的其他组件深度耦合:

7.1 梯度累积(Gradient Accumulation)

累积 K 步才做一次 optimizer.step() 时,等效 batch = K × micro_bs × world_size。此时:

  • 线性缩放律的 k 应取完整等效 batch,而不是单卡 micro_bs。
  • Warmup / 总步数 T_total 应按 optimizer.step() 次数计,而非 forward 次数——否则 scheduler 会「提前衰减」。这是把 scheduler.step() 放在累积循环外层还是内层的经典 bug。

7.2 混合精度(AMP)与 GradScaler


scaler = torch.cuda.amp.GradScaler()
for i, (x, y) in enumerate(loader):
    with torch.cuda.amp.autocast():
        loss = model(x, y) / accum_steps      # 注意:loss 必须除以累积步数
    scaler.scale(loss).backward()
    if (i + 1) % accum_steps == 0:
        scaler.step(optimizer)                  # 只在真正更新时 step
        scaler.update()
        scheduler.step()                        # 与 optimizer.step 配对

陷阱:loss 未除以 accum_steps 会导致梯度被放大 K 倍,等效于隐式把 LR 放大 K 倍——与线性缩放律叠加后会爆炸。

7.3 梯度裁剪(Gradient Clipping)

clip_grad_norm_(1.0) 在 step 前执行。它改变了梯度幅值,因此 LR 的「有效步长」会受 clipping 影响——当梯度范数长期被 clipping 卡在 1.0 时,实际更新量被「上限锁定」,此时再调高 η 收益递减。理解这一点才能正确诊断「为什么调大 LR 没效果」。

8. PyTorch 实战:可复用的组合调度器

下面给出一个生产可用的「Warmup + Cosine/WSD」调度器,支持分布式与梯度累积:


import math
import torch

class WarmupCosineWSD(torch.optim.lr_scheduler._LRScheduler):
    def __init__(self, optimizer, warmup_steps, decay_start, total_steps,
                 peak_lr, min_lr=1e-6, last_epoch=-1):
        self.warmup_steps = warmup_steps
        self.decay_start = decay_start
        self.total_steps = total_steps
        self.peak_lr = peak_lr
        self.min_lr = min_lr
        super().__init__(optimizer, last_epoch)

    def get_lr(self):
        step = self.last_epoch + 1
        if step < self.warmup_steps:
            lr = self.peak_lr * step / max(1, self.warmup_steps)
        elif step < self.decay_start:
            lr = self.peak_lr
        else:
            prog = min((step - self.decay_start) /
                       max(1, self.total_steps - self.decay_start), 1.0)
            lr = self.min_lr + 0.5 * (self.peak_lr - self.min_lr) * \
                 (1 + math.cos(math.pi * prog))
        return [lr for _ in self.base_lrs]

# 用法(分布式 + 梯度累积):
# optimizer = torch.optim.AdamW(model.parameters(), lr=peak_lr, weight_decay=0.1)
# scheduler = WarmupCosineWSD(optimizer, warmup_steps=2000, decay_start=95000,
#                             total_steps=100000, peak_lr=3e-4, min_lr=3e-5)
# 每个 optimizer.step() 之后调用 scheduler.step()

分布式要点:学利率是「逻辑步」概念,应与 world_size 无关——所有 rank 用同一个 step 计数,不要按本地 micro-batch 数各自 step。

9. 生产陷阱清单

# 陷阱 现象 正确做法
1 warmup 步数过短 + 大批量 + 高 LR 训练初期 loss NaN/Inf warmup_steps 随等效 batch 线性放大
2 scheduler.step() 放在累积循环内层 LR 提前衰减到 0 只在 optimizer.step() 时 step
3 loss 未除以 accum_steps 等效 LR 放大 K 倍、发散 backward 前 loss / accum_steps
4 resume 训练未恢复 scheduler 状态 LR 从峰值重来、loss 跳变 保存/加载 scheduler.state_dict()
5 eval/validate 时仍调用 scheduler.step() 验证阶段 LR 被消耗 step 仅发生在 train loop
6 多 optimizer(如 GAN/扩散)共用一个调度 两组参数 LR 同步错误 为每个 optimizer 配独立 scheduler
7 EMA 权重与当前权重混淆 推理用错权重、指标异常 EMA 不进 scheduler,仅平滑推理权重
8 min_lr 设为 0 末期更新量为 0、欠拟合 保留 1e-5~1e-6 余量

10. 结语:与站内系列的方法论桥接

学习率调度是 optimizer 动力学的节流阀,它和本文系列的其他成员构成完整闭环:

  • RMSNorm 解决了每步更新的「数值稳定性底盘」;
  • 相对位置编码 决定了注意力对序列位置的归纳偏置;
  • 交叉注意力 延展了信息跨模态流动的边界;
  • 而 本文的 LR 调度 决定了这些组件在长达数万步的训练中 如何被逐步塑造。

当你在分布式集群上把 batch size 从 1M tokens 扩到 16M tokens 时,记住这条铁律:线性放大 LR、线性放大 warmup、按 optimizer.step 计步、loss 除以累积步数——四者缺一不可,否则再精妙的模型架构也会在训练第一天就 NaN。

下一篇可续:《梯度累积与微批次流水线深度实战》《AdamW 权重解耦与 β 矩估计偏差》,继续补全「Transformer 训练工程」全景。

点赞(0) 打赏

评论列表 共有 0 条评论

暂无评论
立即
投稿
网站二维码

微信公众账号

微信扫一扫加关注

发表
评论
返回
顶部
/* 跳过导航链接 (无障碍) */ .skip-link { position: absolute; top: -100px; left: 15px; z-index: 99999; padding: 8px 16px; background: #007bff; color: #fff; font-size: 14px; border-radius: 0 0 4px 4px; text-decoration: none; transition: top 0.2s; } .skip-link:focus { top: 0; outline: 3px solid #0056b3; }