如何从PyTorch-Lightning的ModelCheckpoint获取最优Checkpoint路径?
如何通过ModelCheckpoint/Trainer内置方法获取最优模型路径?
PyTorch Lightning的ModelCheckpoint回调自带了直接获取最优模型路径的属性,完全不需要手动用glob遍历目录:
- 单最优模型路径:当
save_top_k=1时,直接调用checkpoint_callback.best_model_path,这个属性会返回验证指标最优的模型完整路径。 - Top-K所有模型路径:当
save_top_k>1时,使用checkpoint_callback.best_k_models属性——这是一个字典,键为模型路径,值为对应的验证指标值,且已按指标优劣排序(排序规则由mode参数决定,比如mode="max"时指标越高排名越靠前)。提取字典的键就能得到所有Top-K模型的路径:
# 示例:定义并使用ModelCheckpoint checkpoint_callback = ModelCheckpoint( filename="model_{epoch}-{val_acc:.2f}", save_top_k=3, monitor="val_acc", mode="max" ) trainer = Trainer(callbacks=[checkpoint_callback]) trainer.fit(model) # 获取Top-3模型路径(按val_acc从高到低排序) top_k_paths = list(checkpoint_callback.best_k_models.keys())
补充说明
- 这些属性只有在训练结束后才会被正确赋值,训练过程中访问可能返回空值。
- 如果训练结束后重启程序,需要重新加载回调状态,可以用
ModelCheckpoint.load_checkpoint()方法,加载后依然能访问上述属性:
from lightning.pytorch.callbacks import ModelCheckpoint # 加载已保存的回调状态 checkpoint_callback = ModelCheckpoint.load_checkpoint("/path/to/callback_state.ckpt") top_k_paths = list(checkpoint_callback.best_k_models.keys())
内容的提问来源于stack exchange,提问作者Daraan
相关产品推荐
相关产品推荐

