You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

PyTorch Lightning 1.3.8两阶段训练如何重置学习率?

解决方案

针对你在PyTorch Lightning 1.3.8 + Torch 1.9.0下的两阶段训练学习率重置问题,这里提供两种可行的方案:

方案1:仅加载模型权重,重新初始化优化器与调度器

绕过加载checkpoint中的优化器和调度器状态,手动加载模型权重后重新构建优化器和调度器,完全从第二阶段的初始学习率开始训练:

# 第二阶段实例化模型,关闭teacher forcing
model = YourModel(use_teacher_forcing=False)

# 加载第一阶段的训练checkpoint
checkpoint = torch.load("path/to/first_stage_checkpoint.ckpt")
# 仅加载模型的参数权重,忽略优化器、调度器状态
model.load_state_dict(checkpoint["state_dict"])

# 重新初始化优化器,设置第二阶段初始学习率1e-4
optimizer = torch.optim.Adam(model.parameters(), lr=1e-4)

# 定义适配第二阶段的学习率调度器
# 全局epoch里程碑设为7、9(对应第二阶段第2、4轮,匹配1e-5、1e-6的衰减节点)
scheduler = torch.optim.lr_scheduler.MultiStepLR(
    optimizer,
    milestones=[7, 9],
    gamma=0.1
)

# 初始化Trainer并启动训练
trainer = pl.Trainer(max_epochs=9, ...)
trainer.fit(model, optimizer=optimizer, lr_scheduler=scheduler)

这种方案彻底避免了旧调度器状态的干扰,逻辑清晰,适合不需要保留第一阶段优化器动量等状态的场景。

方案2:加载完整checkpoint后手动重置调度器与学习率

如果需要保留第一阶段优化器的状态(比如动量),可以在训练开始时通过Lightning的钩子手动重置学习率和调度器参数:

class YourModel(pl.LightningModule):
    def __init__(self, use_teacher_forcing):
        super().__init__()
        self.use_teacher_forcing = use_teacher_forcing
        # 模型层初始化...

    def configure_optimizers(self):
        optimizer = torch.optim.Adam(self.parameters(), lr=1e-4)
        # 这里先按通用框架定义调度器,后续在钩子中调整
        scheduler = torch.optim.lr_scheduler.MultiStepLR(
            optimizer,
            milestones=[2, 4],
            gamma=0.1
        )
        return {"optimizer": optimizer, "lr_scheduler": scheduler}

    def on_train_start(self):
        # 仅在第二阶段(关闭teacher forcing)执行重置逻辑
        if not self.use_teacher_forcing:
            # 重置优化器学习率为1e-4
            optimizer = self.optimizers()
            for param_group in optimizer.param_groups:
                param_group["lr"] = 1e-4
            
            # 重置调度器状态,重新计数并更新里程碑
            scheduler = self.lr_schedulers()
            # 将last_epoch设为-1,让调度器从初始状态开始计数
            scheduler.last_epoch = -1
            # 更新里程碑为第二阶段的全局epoch节点
            scheduler.milestones = {7: 0.1, 9: 0.1}

第二阶段启动训练时,直接用resume_from_checkpoint加载第一阶段的checkpoint:

model = YourModel(use_teacher_forcing=False)
trainer = pl.Trainer(resume_from_checkpoint="path/to/first_stage_checkpoint.ckpt", max_epochs=9, ...)
trainer.fit(model)

注意事项

  • 确保PyTorch Lightning 1.3.8中,self.optimizers()和self.lr_schedulers()能正确获取到对应的对象(早期版本可能需要用self.trainer.optimizers[0]代替)。
  • MultiStepLR的milestones属性在1.9.0中是公开的字典类型,可以直接修改;如果使用其他调度器,需要对应调整状态重置的逻辑。

内容的提问来源于stack exchange,提问作者M_M246

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.08.14 01:50:24