PyTorch Lightning多标签文本分类训练:ModelCheckpoint参数缺失报错解决
解决PyTorch Lightning ModelCheckpoint的TypeError问题
你的报错是因为使用的PyTorch Lightning版本较低,旧版本的ModelCheckpoint类初始化需要filepath参数,而非你当前用的dirpath+filename组合参数。下面给出两种解决方式:
方式一:适配旧版本参数
将dirpath和filename合并为filepath参数,直接指定完整的checkpoint文件路径:
checkpoint_callback = ModelCheckpoint( filepath="checkpoints/best-checkpoint", save_top_k=1, verbose=True, monitor="val_loss", mode="min" )
如果需要动态命名(比如包含epoch、指标值),可以在filepath中使用占位符:
checkpoint_callback = ModelCheckpoint( filepath="checkpoints/best-checkpoint-{epoch:02d}-{val_loss:.2f}", save_top_k=1, verbose=True, monitor="val_loss", mode="min" )
方式二:升级PyTorch Lightning到新版本
如果你想保留dirpath和filename的写法,直接升级PyTorch Lightning到1.0及以上版本,新版本已经废弃了filepath,改用dirpath(指定保存目录)和filename(指定文件名模板)的组合。执行以下命令升级:
pip install --upgrade pytorch-lightning
升级后你原来的代码就可以正常运行了。
内容的提问来源于stack exchange,提问作者joe
相关产品推荐
相关产品推荐

