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

如何在PyTorch Lightning中控制trainer.global_step,避免多优化器计数翻倍?

解决PyTorch Lightning多优化器导致global_step翻倍的问题

当关闭自动优化(self.automatic_optimization = False)时,PyTorch Lightning默认会在每次调用optimizer.step()后递增global_step,这会导致单训练步骤使用多优化器时步数计数异常。以下是简便的解决方法:

核心思路

手动控制global_step的递增逻辑:阻止每个优化器的step()自动触发步数增长,在所有优化器执行完成后统一更新一次global_step。

具体实现

在你的LightningModule的training_step方法中,调用优化器的step()时传入update_global_step=False参数,最后手动调用reduce_global_step(1)来完成一次步数递增:

def training_step(self, batch, batch_idx):
    # 获取多个优化器
    opt1, opt2 = self.optimizers()

    # 处理第一个优化器的梯度更新
    loss1 = self.calculate_loss1(batch)
    self.manual_backward(loss1)
    opt1.step(update_global_step=False)  # 不自动更新global_step
    opt1.zero_grad()

    # 处理第二个优化器的梯度更新
    loss2 = self.calculate_loss2(batch)
    self.manual_backward(loss2)
    opt2.step(update_global_step=False)  # 不自动更新global_step
    opt2.zero_grad()

    # 手动统一更新一次global_step
    self.trainer.strategy.reduce_global_step(1)

补充说明

  • 该方法保持了global_step与单优化器训练时的计数逻辑一致,确保检查点回调的触发时机、保存的步数标签完全对齐。
  • 若使用学习率调度器,无需额外调整:调度器默认依赖global_step,手动更新后的步数会自动适配调度逻辑。
  • 避免直接修改trainer.global_step(只读属性),通过官方提供的reduce_global_step方法更新是符合框架设计的安全方式。

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.17 14:01:03