如何在PyTorchLightning中手动指定Checkpoint的存储路径与文件名?
手动指定TensorBoardLogger关联的Checkpoint存储路径与文件名
TensorBoardLogger仅负责日志记录,Checkpoint的存储规则由ModelCheckpoint回调单独控制,你只需自定义该回调的参数,就能实现手动指定存储文件夹与文件名,具体操作如下:
1. 导入必要模块
from pytorch_lightning import Trainer from pytorch_lightning.loggers import TensorBoardLogger from pytorch_lightning.callbacks import ModelCheckpoint
2. 配置自定义ModelCheckpoint回调
通过dirpath指定存储文件夹,filename自定义文件名格式(支持动态占位符):
# 自定义Checkpoint存储规则 checkpoint_callback = ModelCheckpoint( dirpath="./my_custom_checkpoints", # 手动设置存储文件夹,可填绝对/相对路径 filename="model_epoch_{epoch}_val_acc_{val_acc:.2f}", # 自定义文件名,支持占位符 save_top_k=3, # 保存性能Top3的模型 monitor="val_acc", # 监控的验证指标 mode="max" # 指标最大化时判定为更优 )
dirpath:目标文件夹不存在时,PyTorch Lightning会自动创建filename占位符:支持{epoch}(当前训练轮数)、{global_step}(全局步数)、{train_loss}(训练损失)等所有训练过程中记录的变量
3. 关联Logger与Trainer
将自定义的checkpoint_callback和TensorBoardLogger一起传入Trainer即可:
# 初始化TensorBoardLogger logger = TensorBoardLogger(save_dir="./tb_logs", name="my_experiment") # 启动训练 trainer = Trainer( max_epochs=20, logger=logger, callbacks=[checkpoint_callback] ) trainer.fit(your_model)
可选:将Checkpoint存到TensorBoard日志目录下
如果希望Checkpoint与TensorBoard日志放在同一父目录,可利用logger.log_dir获取日志的实际存储路径:
logger = TensorBoardLogger(save_dir="./tb_logs", name="my_experiment") checkpoint_callback = ModelCheckpoint( dirpath=f"{logger.log_dir}/custom_checkpoints", # 存到日志目录的子文件夹中 filename="best_model_{epoch}_{val_loss:.3f}", monitor="val_loss", mode="min" )
额外提示
- 若想每次保存都覆盖同一个文件,直接设置
filename="latest_model"即可 - 需要保存所有Checkpoint时,设置
save_top_k=-1,文件名中的占位符会自动避免文件覆盖
内容的提问来源于stack exchange,提问作者LukeTheWalker
相关产品推荐
相关产品推荐

