使用Pytorch Lightning传递ModelCheckpoint时出现配置错误如何解决?
问题修复方案
报错原因
- 你所使用的PyTorch Lightning版本中,
pl.Trainer的checkpoint_callback参数仅支持传入布尔值(True/False),用于控制是否启用默认 checkpoint 功能,不支持直接传入自定义的ModelCheckpoint实例。 - 所有自定义回调(包括早停、自定义 checkpoint 等)都需要统一传入
callbacks参数的列表中,这是当前版本的强制要求。
修复代码
只需要调整pl.Trainer的初始化参数,删除单独的checkpoint_callback传参,把自定义的 checkpoint 回调实例加入callbacks列表即可:
trainer = pl.Trainer( logger=logger, callbacks=[early_stopping_callback, checkpoint_callback], max_epochs=N_EPOCHS, gpus=1, progress_bar_refresh_rate=30 )
补充说明
修改后你的训练逻辑会完全保留原有设定规则:连续2轮验证集损失不下降就触发早停,同时仅保存验证集损失最低的1份权重到指定路径。
内容的提问来源于stack exchange,提问作者MAC
相关产品推荐
相关产品推荐

