PyTorch Lightning TypeError:__init__()收到意外关键字参数'checkpoint_callback'
问题与解决方案
错误信息
TypeError Traceback (most recent call last) <ipython-input-41-2892cdd4e738> in <module>() 5 max_epochs=N_EPOCHS, 6 gpus=1, #GPU ----> 7 progress_bar_refresh_rate=30 8 ) /usr/local/lib/python3.7/dist-packages/pytorch_lightning/utilities/argparse.py in insert_env_defaults(self, *args, **kwargs) 343 344 # all args were already moved to kwargs --> 345 return fn(self, **kwargs) 346 347 return cast(_T, insert_env_defaults) TypeError: __init__() got an unexpected keyword argument 'checkpoint_callback'
运行的代码片段
trainer = pl.Trainer( logger=logger, checkpoint_callback=checkpoint_callback, callbacks=[early_stopping_callback], max_epochs=N_EPOCHS, gpus=1, #GPU progress_bar_refresh_rate=30 )
checkpoint_callback定义
checkpoint_callback = ModelCheckpoint( dirpath="checkpoints", filename="best-checkpoint", save_top_k=1, verbose=True, monitor="val_loss", mode="min" )
解决方案
错误原因是PyTorch Lightning 1.5及以上版本中,checkpoint_callback不再是pl.Trainer的独立参数,所有回调类(包括模型检查点、早停)都需要统一放入callbacks列表传递。
修改后的Trainer初始化代码:
trainer = pl.Trainer( logger=logger, callbacks=[early_stopping_callback, checkpoint_callback], max_epochs=N_EPOCHS, gpus=1, #GPU progress_bar_refresh_rate=30 )
只需移除checkpoint_callback=checkpoint_callback这一行,将checkpoint_callback实例添加到callbacks列表中即可解决该错误。
内容的提问来源于stack exchange,提问作者Quantizer
相关产品推荐
相关产品推荐

