PyTorch中学习率预热后延迟启动余弦退火重启调度器方案问询
解决PyTorch预热+学习率保持+余弦退火重启的调度问题
核心方案思路
把训练流程拆成三个独立阶段,通过全局步数精准控制每个阶段的学习率逻辑,彻底避免调度器在非目标阶段干扰学习率:
- 预热阶段:前N步线性提升学习率至预设最大值
- 保持阶段:接下来M步维持最大学习率不变
- 调度阶段:剩余步数启动余弦退火重启策略
修改后代码实现
import torch import pytorch_warmup as warmup # 超参数配置 lr = 1e-3 # 预热后的目标学习率(max_lr) warmup_steps = 1000 # 预热总步数 hold_steps = 500 # 保持max_lr的步数 T_0 = 10 # 余弦退火重启的初始周期(单位:epoch) T_mult = 2 num_epochs = ... # 你的总训练epoch数 dataloader = ... # 你的数据加载器 params = ... # 模型参数 optimizer = torch.optim.AdamW(params, lr=lr) # 初始化余弦退火重启调度器:初始学习率设为max_lr,确保调度启动时从该值开始衰减 lr_scheduler = torch.optim.lr_scheduler.CosineAnnealingWarmRestarts(optimizer, T_0=T_0, T_mult=T_mult) # 初始化线性预热调度器:从0线性提升至optimizer的初始lr(即max_lr) warmup_scheduler = warmup.LinearWarmup(optimizer, warmup_period=warmup_steps) iters_per_epoch = len(dataloader) global_step = 0 start_sched_steps = warmup_steps + hold_steps # 调度器启动的起始步数 for epoch in range(1, num_epochs + 1): for idx, batch in enumerate(dataloader): optimizer.zero_grad() # 前向传播与损失计算 loss = ... loss.backward() optimizer.step() # 分阶段处理学习率调整 if global_step < warmup_steps: # 预热阶段:仅应用线性预热,不启动余弦调度器 warmup_scheduler.step() elif global_step < start_sched_steps: # 保持阶段:不做任何学习率调整,维持当前max_lr pass else: # 调度阶段:计算调度器的相对起始epoch,保证周期计算不受前面阶段影响 relative_epoch = (global_step - start_sched_steps) / iters_per_epoch lr_scheduler.step(relative_epoch) global_step += 1
关键细节说明
- 阶段隔离:通过
global_step严格划分三个阶段,确保余弦调度器仅在预热+保持阶段结束后才启动,彻底解决预热阶段被调度器干扰的问题 - 预热逻辑:单独调用
warmup_scheduler.step(),保证学习率纯线性上升至max_lr,避免原代码中预热因子与调度器LR相乘的异常 - 调度器起始校准:计算
relative_epoch让余弦调度器从第0个周期开始计算,T_0的周期设置完全符合预期,不需要因为预热阶段而调大T_0值 - 保持阶段灵活性:
hold_steps可根据需求设为0,即预热后直接启动衰减调度
内容的提问来源于stack exchange,提问作者Molem7b5
相关产品推荐
相关产品推荐

