学习率调度深度实战:从 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),但存在两个硬伤:
- 突变不连续:LR 在断点处瞬间跳变,优化轨迹会在断点附近产生可见的 loss 台阶(见图 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) 把训练切成三段:
- Warmup:爬升到
η_max(同前)。 - Stable(恒定段):在
η_max恒定训练绝大部分步数——这是算力利用率最高的阶段。 - 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 训练工程」全景。

发表评论 取消回复