为何PyTorch Lightning的global_step与WandB绘图步长不同步?
问题解答
1. 原因分析
PyTorch Lightning中的self.global_step是全局步数计数器,它会累加所有类型的步骤——包括训练阶段的step、验证阶段的step(甚至测试阶段的step)。默认情况下,Trainer在每个训练epoch结束后会自动启动验证流程,验证过程中的每个step都会让global_step递增。因此你在training_step中记录的global_step,其实已经包含了之前验证阶段产生的步数,最终导致WandB图表里的train_step看起来混合了训练和验证的步数。
2. 解决方法
有两种简单可靠的方式实现只记录训练步骤:
方式一:使用训练循环专属的步数计数器
直接调用PyTorch Lightning训练循环的专属步数,它不会被验证/测试步骤干扰:
wandb.log({"train_step": self.trainer.train_loop.global_step})
方式二:手动维护训练步数计数器
在你的LightningModule类中初始化一个自定义计数器,仅在训练步骤中递增:
class YourModel(pl.LightningModule): def __init__(self): super().__init__() self.train_step_counter = 0 # 初始化训练步数计数器 def training_step(self, batch, batch_idx): # ... 你的训练逻辑 ... self.train_step_counter += 1 wandb.log({"train_step": self.train_step_counter}) # ... 返回损失 ...
这样就能确保WandB中记录的train_step只对应训练阶段的步骤。
内容的提问来源于stack exchange,提问作者user3668129
相关产品推荐
相关产品推荐

