如何在PyTorch Lightning中记录每次保存的Checkpoint路径
解决方案
方法1:自定义带日志的Checkpoint回调
直接继承ModelCheckpoint,重写核心保存方法,在Checkpoint完成保存后立即记录路径:
from pytorch_lightning.callbacks import ModelCheckpoint import logging logger = logging.getLogger(__name__) class LoggingModelCheckpoint(ModelCheckpoint): def _save_checkpoint(self, trainer, filepath): # 先执行父类的保存逻辑 super()._save_checkpoint(trainer, filepath) # 记录Checkpoint完整路径 logger.info(f"Checkpoint saved to: {filepath}") # 替换原Checkpoint回调 checkpoint_callback = LoggingModelCheckpoint( filename="fa_classifier_{epoch:02d}", every_n_epochs=val_every_n_epochs, save_top_k=-1, ) # Trainer初始化保持不变 trainer = Trainer( callbacks=[checkpoint_callback], default_root_dir=checkpoints_path, check_val_every_n_epoch=val_every_n_epochs, max_epochs=max_epochs, gpus=1 )
方法2:新增独立回调读取已保存路径
如果不想修改原Checkpoint逻辑,可以添加一个单独的回调,在验证阶段结束后读取已保存的Checkpoint路径:
from pytorch_lightning.callbacks import Callback, ModelCheckpoint import logging logger = logging.getLogger(__name__) class CheckpointPathLogger(Callback): def on_validation_end(self, trainer, pl_module): # 找到Trainer中的Checkpoint回调实例 checkpoint_cb = next(cb for cb in trainer.callbacks if isinstance(cb, ModelCheckpoint)) # 读取最新保存的Checkpoint路径 if checkpoint_cb.saved_checkpoints: latest_path = checkpoint_cb.saved_checkpoints[-1] logger.info(f"Latest checkpoint saved to: {latest_path}") # 初始化原Checkpoint回调和日志回调 checkpoint_callback = ModelCheckpoint( filename="fa_classifier_{epoch:02d}", every_n_epochs=val_every_n_epochs, save_top_k=-1, ) path_logger = CheckpointPathLogger() # 同时传入两个回调到Trainer trainer = Trainer( callbacks=[checkpoint_callback, path_logger], default_root_dir=checkpoints_path, check_val_every_n_epoch=val_every_n_epochs, max_epochs=max_epochs, gpus=1 )
说明
- 方法1的时机最精准,在Checkpoint完成写入后立即触发日志,不会有延迟。
- 方法2通过监听验证结束事件,利用
ModelCheckpoint自带的saved_checkpoints属性获取历史保存路径,适合无需改动原Checkpoint逻辑的场景。
内容的提问来源于stack exchange,提问作者Gulzar
相关产品推荐
相关产品推荐

