如何在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
相关产品推荐
相关产品推荐

