PyTorch CosineAnnealingLR的T_max和eta_min参数设置

在深度学习训练中,学习率调度器(Learning Rate Scheduler)是至关重要的超参数优化工具。其中,余弦退火(Cosine Annealing)因其平滑的衰减特性和优秀的性能表现,已成为最流行的调度策略之一。

CosineAnnealingLR 基础

余弦退火的学习率变化遵循余弦函数:

\(\eta_t=\eta_{min}+\frac12(\eta_{max}-\eta_{min})\Big(1+\cos\big(\frac{t}{T_{max}}\pi\big)\Big)\)

  • \(\eta_t\):当前步学习率
  • \(\eta_{max}\):优化器初始学习率
  • \(\eta_{min}\):eta_min,学习率下限
  • t:当前 step(batch 迭代次数)
  • \(T_{max}\):余弦周期长度

当 \(t=T_{max}\),\(\cos(\pi)=-1\),学习率 = \(\eta_{min}\)。

基本用法

python 复制代码
import torch
import torch.optim as optim
from torch.optim.lr_scheduler import CosineAnnealingLR

# 定义优化器
optimizer = optim.Adam(model.parameters(), lr=0.001)

# 创建余弦退火调度器
scheduler = CosineAnnealingLR(
    optimizer,
    T_max=100,      # 周期长度
    eta_min=1e-6    # 最小学习率
)

T_max 参数

T_max 控制学习率从最大值下降到最小值所需的步数。

整个训练全部迭代步数,一次完整余弦衰减,训练结束 lr 到达 eta_min。

通常设置为:

  1. T_max = total_iter = (样本数 // batch_size) × epochs
  2. T_max = epochs
配置 调用 scheduler.step () 时机
T_max = total_iterations 每个 batch 内部执行一次(每 1 个 step 调用)
T_max = epochs epoch 循环末尾调用,一个 epoch 只 step 一次

T_max = total_iterations(总迭代 step)

使用条件(两个必须同时满足)

scheduler.step() 写在内层 batch 循环,每跑完一个 batch 调用一次

python 复制代码
for epoch in range(epochs):
    for imgs, label in dataloader:
        # 前向、loss、反向传播
        optimizer.step()
        scheduler.step() #每个batch调用

T_max = total_iterations = (num_samples // batch_size) * epochs

T_max = epochs(T_max 数值等于轮数)

调度器调用时机

python 复制代码
# 正确配套写法
scheduler = CosineAnnealingLR(optimizer, T_max=epochs, eta_min=1e‑6)

for epoch in range(epochs):
    for imgs, label in dataloader:
        #训练,只更新optimizer,不调用scheduler.step()
        optimizer.step()
    scheduler.step() #❗每个epoch结束才调用一次

假设epoch=100,学习率初始值为0.0006,

当T_max = epochs,则学习率变化;

当T_max < epochs,则学习率变化;

当T_max > epochs,则学习率变化;下图学习率最小值停留在0.0003

eta_min 参数

eta_min:余弦退火的学习率最低下限。

eta_min=0:训练末尾学习率直接归零,参数不再更新;

eta_min=1e‑6(工程常用):保留极小学习率,允许参数微小更新,避免训练末期完全冻结。

不建议直接置 0,部分场景会出现收敛停滞。CV 任务(车道线、检测)推荐eta_min=1e‑6 ~ 1e‑7。

相关推荐
夏天拐跑了西瓜8 小时前
一文入门LangChain:从框架认知到构建你的第一个AI Agent
python·langchain·conda·agent
奕鼎竜瑆8 小时前
Solid 前端响应式开发从零到精通
前端·人工智能
YOLO数据集集合8 小时前
无人机桥梁损伤目标检测数据集 | 桥梁损伤 无人机巡检 结构健康监测 多类别检测9135期
人工智能·目标检测·无人机
动恰客流统计8 小时前
景区客流统计怎么做?兼顾管控与运营的实施方案解析
大数据·前端·人工智能
weixin_443883018 小时前
合规整改倒计时:高等级签名证书的应用场景
大数据·人工智能·法大大·法大大电子签·电子合同
魔猴疯猿8 小时前
从0到1用Python开发第一个智能体
人工智能·python·深度学习·神经网络·机器学习
天一生水water9 小时前
变分模态分解(VMD)的教程
人工智能
荆棘鸟智能9 小时前
林业遥感长势评估怎么做?从NDVI时序分析到LiDAR蓄积量回归
人工智能·算法·智慧林业
代码方舟9 小时前
数据科学风控实战:基于天远全能消金报告构建自动化信用评估网关
运维·人工智能·自动化