如何禁止PyTorch Lightning中TensorBoard logger默认记录当前epoch
解决方案
自动记录的epoch指标是PyTorch Lightning的Trainer在上报日志时默认塞入全局metrics字典的,不属于TensorBoardLogger的单独显式逻辑,因此你在logger源码里找不到对应的记录调用,目前官方没有提供直接关闭该记录的配置项,可通过自定义Logger子类的方式实现过滤:
操作步骤
- 自定义继承自TensorBoardLogger的子类,重写log_metrics方法过滤epoch键:
from pytorch_lightning.loggers import TensorBoardLogger class NoEpochTensorBoardLogger(TensorBoardLogger): def log_metrics(self, metrics, step=None): filtered_metrics = {k: v for k, v in metrics.items() if k != "epoch"} super().log_metrics(filtered_metrics, step=step)
- 初始化Trainer时替换默认的TensorBoardLogger即可,搭配你已使用的
default_hp_metric=False参数可完全清除默认冗余记录:
logger = NoEpochTensorBoardLogger( save_dir="./logs", default_hp_metric=False ) trainer = Trainer( logger=logger, # 其他Trainer配置项 )
该方法兼容所有PyTorch Lightning版本,不需要修改框架源码,也不会影响其他自定义指标的正常记录。

内容的提问来源于stack exchange,提问作者Paul Hager
相关产品推荐
相关产品推荐

