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

相关推荐
aichitang20241 小时前
快乐泛函每一天!内积空间
c++·python·数学·算法·机器学习·ai·泛函分析
环境栈笔记1 小时前
指纹浏览器怎么用:从 Profile、代理到环境检测的完整上手流程
前端·人工智能·后端·自动化
光锥智能1 小时前
他山科技亮相WRC 2026:“机器人幼儿园”驱动具身智能迈向“经验时代”
人工智能·科技·机器人
Raas1001 小时前
AI网关能省多少钱?MAI Gateway (魔芋企业级AI网关)降本ROI实战案例
人工智能
jikemaoshiyanshi1 小时前
自建大模型推理服务如何优化算力成本与资源效率?—— 基于智能路由、PD 分离、缓存、弹性调度的云上基建选型
人工智能
土拨鼠不是老鼠1 小时前
python 使用 plotly 进行数据分析
python·plotly·数据分析
江畔柳前堤1 小时前
LLM + Agent 模型效果评估:从入门到工业级体系构建的完整指南
开发语言·人工智能·自然语言处理·chatgpt·架构·json·batch
Elastic 中国社区官方博客2 小时前
跳过 mapping 爆炸:ES|QL 无需动态 mapping 即可查询无 schema JSON key
大数据·人工智能·sql·elasticsearch·搜索引擎·json·全文检索
欧特克_Glodon2 小时前
OpenCV计算机视觉开发入门与实践<十四>:基础几何图形绘制
人工智能·opencv·计算机视觉