不能直接改 param_groups0,因为调度器依赖 last_epoch 和缓存 lr 值同步步进逻辑;手动修改不更新 last_epoch,导致 step() 覆盖预期值,引发学习率跳变、get_last_lr() 与实际不一致、恢复训练时重启初始值等问题。

PyTorch 的 torch.optim.lr_scheduler 支持自定义调度器,但必须继承 _LRScheduler(或其子类如 LambdaLR),且需重写 get_lr();直接修改 optimizer.param_groups 中的 lr 会导致调度器状态错乱。
为什么不能直接改 param_groups[0]['lr']
PyTorch 调度器内部依赖 self.last_epoch 和缓存的 epoch 级 lr 值来同步步进逻辑。手动改 lr 不更新 last_epoch,下次调用 scheduler.step() 会覆盖或跳过预期值,尤其在 step() 频率与训练循环不一致时(比如每 batch step 一次但 scheduler 设计为每 epoch step)。
常见错误现象:
- 学习率突然归零或跳变
-
lr_scheduler.get_last_lr()返回值与实际param_groups[0]['lr']不一致 - 恢复训练时 lr 从初始值重启,而非接续上一个 epoch
最简可行:继承 _LRScheduler 并实现 get_lr()
这是最可控、兼容性最好的方式。注意:_LRScheduler 是私有基类,但官方文档明确允许继承它来自定义逻辑。
立即学习“Python免费学习笔记(深入)”;
实操建议:
- 必须在
__init__中调用super().__init__(optimizer, last_epoch) -
get_lr()必须返回一个list,长度等于optimizer.param_groups数量,每个元素是对应组的学习率 - 若只调第一组,其他组保持原 lr,可写成
[new_lr] + [group['lr'] for group in self.optimizer.param_groups[1:]] - 避免在
get_lr()中做耗时计算(如读文件、网络请求),它会在每次step()时被调用
示例:线性预热 + 余弦衰减组合调度器
from torch.optim.lr_scheduler import _LRScheduler
import math
<p>class LinearWarmupCosineAnnealing(_LRScheduler):
def <strong>init</strong>(self, optimizer, warmup_epochs, max_epochs, eta_min=0, last_epoch=-1):
self.warmup_epochs = warmup_epochs
self.max_epochs = max_epochs
self.eta_min = eta_min
super().<strong>init</strong>(optimizer, last_epoch)</p><pre class="brush:php;toolbar:false;">def get_lr(self):
if self.last_epoch < self.warmup_epochs:
# 线性预热:lr 从 0 增至 base_lr
return [
base_lr * self.last_epoch / max(1, self.warmup_epochs)
for base_lr in self.base_lrs
]
else:
# 余弦衰减:从 base_lr 到 eta_min
t = self.last_epoch - self.warmup_epochs
T = self.max_epochs - self.warmup_epochs
return [
self.eta_min + 0.5 * (base_lr - self.eta_min) * (1 + math.cos(math.pi * t / T))
for base_lr in self.base_lrs
]更轻量替代:用 LambdaLR 匿名函数快速实验
适合临时验证调度逻辑,无需定义新类。但注意:LambdaLR 的 lr_lambda 接收的是 int 类型的 epoch(不是 last_epoch),且默认按 epoch 步进;若需 per-batch 调度,得配合 StepLR 或手动传入迭代序号。
使用场景:
- 快速试跑不同衰减曲线(如指数、多项式)
- 调试时绕过完整类定义
- 多参数组用不同 lambda(传入 list of lambdas)
示例:每 10 个 epoch 将 lr 乘以 0.9
lambda epoch: 0.9 ** (epoch // 10) scheduler = LambdaLR(optimizer, lr_lambda=lambda epoch: 0.9 ** (epoch // 10))
⚠️ 注意:该 lambda 在 step() 时被调用,传入的是当前 last_epoch + 1,所以实际生效从 epoch=1 开始;若想 epoch=0 就应用,需在 lambda 内部处理偏移。
容易被忽略的细节:step() 调用时机与 last_epoch 初始化
调度器的 last_epoch 默认为 -1,首次 step() 后变为 0。这意味着如果你在训练前先 step() 一次(比如为了 warmup),必须确保 last_epoch 已正确设置,否则 get_lr() 会基于 -1 计算出异常值(如负数 lr)。
关键点:
- 加载 checkpoint 时,必须同时加载
scheduler.state_dict()并调用scheduler.load_state_dict() - 若 checkpoint 没保存 scheduler 状态,需手动设
scheduler.last_epoch = start_epoch(不能只改last_epoch,还要确保base_lrs与当前optimizer.param_groups一致) - per-batch 调度时,
step()应放在每个 batch 后,但LambdaLR默认按 epoch 计数 —— 此时应改用MultiplicativeLR或自定义类,并把last_epoch当作全局 step 计数器用
复杂点在于:调度器行为高度耦合于你如何组织训练循环(epoch-based vs step-based)、是否 resume、以及是否 multi-stage。没对齐这些,再精巧的公式也白搭。


















