“线性预热 + 余弦衰减”学习率策略
收藏回复举报
“线性预热 + 余弦衰减”学习率策略
发表于2025-10-25 23:06:10
0 查看

在深度学习训练中,学习率调度(Learning Rate Scheduling)是影响模型收敛速度与最终性能的关键因素。尽管现代优化器(如AdamW)具备自适应能力,但合理的学习率变化策略仍不可或缺。

在BERT、ViT、LLaMA等主流大模型的训练配置中,我们几乎总能看到这样一行代码:

scheduler = get_linear_schedule_with_warmup(optimizer, num_warmup_steps=10000, num_training_steps=1000000)

这背后隐藏着一种被广泛验证有效的学习率策略组合:线性预热(Linear Warmup) + 余弦衰减(Cosine Decay)。


一、训练初期的“梯度爆炸”风险

在模型训练的第一步,参数是随机初始化的,网络尚未建立任何语义理解。此时:

  • 前向传播的输出分布极不均匀;
  • 反向传播计算出的梯度方差极大;
  • 若此时使用全量学习率更新参数,会导致:
    • 损失函数剧烈震荡;
    • 参数跳入“坏”的局部最优;
    • 甚至梯度溢出(NaN)导致训练失败。

典型现象:

训练开始后Loss从8.0跳到9.5,再跳回7.8,反复震荡,迟迟无法下降。


二、第一阶段:线性预热(Linear Warmup)

目标:平稳启动

在前 N 个训练步(如10,000步)内,学习率从 0 线性增长到目标值(如 5e-4):

$$ \text{lr}(t) = \text{lr}_{\text{base}} \times \frac{t}{N}, \quad t \in [0, N] $$

为什么是“线性”?

  • 简单可控:增长速率恒定,易于实现和调试;
  • 避免激进更新:小学习率限制了初始参数更新的“步长”,防止模型在混乱状态下走偏;
  • 配合BN层收敛:Batch Normalization需要多个batch统计均值和方差,预热期间为其提供稳定环境。

 经验法则:预热步数通常占总训练步数的 1%~10%。例如:

  • BERT Base:10,000步预热(占总步数1%)
  • ViT-L/16:5,000步预热

三、第二阶段:余弦衰减(Cosine Annealing)

预热结束后,学习率不再保持恒定,而是按余弦函数缓慢下降至接近0:

$$ \text{lr}(t) = \text{lr}{\text{min}} + \frac{1}{2} (\text{lr}{\text{base}} - \text{lr}_{\text{min}}) \left(1 + \cos\left(\pi \cdot \frac{t - N}{T - N}\right)\right) $$

其中:

  • $ N $:预热结束步数
  • $ T $:总训练步数
  • $ \text{lr}_{\text{min}} $:最小学习率(常设为0)

为什么是“余弦”而非“线性”?

特性余弦衰减线性衰减
下降速度前快后慢匀速下降
收敛行为末期小学习率精细调参可能过早收敛
泛化性能更好,避免陷入尖锐极小一般
实践效果被大模型广泛采用逐渐被替代

关键优势:

余弦衰减在训练后期提供极小的学习率,让模型在损失曲面的平坦区域进行精细搜索,提升泛化能力。


四、为何“线性+余弦”是黄金组合?

阶段策略作用
0 ~ Warmup Steps线性增长防止初期震荡,稳定启动
Warmup ~ Total Steps余弦下降平稳收敛,提升泛化

这种组合实现了**“稳启动 + 慢收敛”** 的理想训练动态:

  1. 前期:像“热身运动”,让模型适应数据分布;
  2. 后期:像“精细打磨”,在最优解附近微调参数。

五、代码实现(Hugging Face Transformers 风格)

from transformers import get_cosine_schedule_with_warmup
from torch.optim import AdamW

optimizer = AdamW(model.parameters(), lr=5e-4)
total_steps = 1_000_000
warmup_steps = 10_000

# 构建调度器
scheduler = get_cosine_schedule_with_warmup(
    optimizer,
    num_warmup_steps=warmup_steps,
    num_training_steps=total_steps
)

# 训练循环中每步调用
for step, batch in enumerate(dataloader):
    loss = model(batch).loss
    loss.backward()
    optimizer.step()
    scheduler.step()  # 更新学习率

 

🔧 调参建议:

  • warmup_steps:从 1000 开始尝试,观察Loss是否平稳下降;
  • lr_min:通常设为 0 或 1e-7;
  • 若训练后期Loss卡住,可尝试延长余弦阶段。

“预热是让模型‘冷静下来’,余弦衰减是让它‘沉下心来’——两者结合,方能训出好模型。”

 

本帖最后由 匿名用户 于 2026/09/03 15:10:58 编辑

我要发帖子