PyTorch Lightning加载Checkpoint报错:存在意外键bert.embeddings.position_ids
解决PyTorch Lightning加载BERT模型Checkpoint时"bert.embeddings.position_ids"错误
这个错误是因为BERT模型中的position_ids是动态生成的非可训练参数,却被意外存入了Checkpoint,但加载时模型的state_dict中没有这个key,导致参数不匹配。以下是几种解决方法:
方法1:加载时过滤多余的key
先加载Checkpoint文件,删除不需要的position_ids键后再加载到模型:
import torch import pytorch_lightning as pl # 加载Checkpoint checkpoint = torch.load("example.ckpt") # 删除冗余键 del checkpoint["state_dict"]["bert.embeddings.position_ids"] # 初始化模型并加载处理后的state_dict model = QTagClassifier() model.load_state_dict(checkpoint["state_dict"])
如果使用PyTorch Lightning的load_from_checkpoint,可以在模型类中重写on_load_checkpoint方法自动处理:
class QTagClassifier(pl.LightningModule): # 模型原有代码... def on_load_checkpoint(self, checkpoint): # 移除Checkpoint中的冗余键 if "bert.embeddings.position_ids" in checkpoint["state_dict"]: del checkpoint["state_dict"]["bert.embeddings.position_ids"]
之后直接调用model = QTagClassifier.load_from_checkpoint("example.ckpt")即可正常加载。
方法2:保存时排除非训练参数
在模型类中重写state_dict方法,保存时自动过滤掉position_ids:
class QTagClassifier(pl.LightningModule): # 模型原有代码... def state_dict(self, *args, **kwargs): original_state_dict = super().state_dict(*args, **kwargs) # 移除非训练的position_ids if "bert.embeddings.position_ids" in original_state_dict: del original_state_dict["bert.embeddings.position_ids"] return original_state_dict
这样后续保存的Checkpoint就不会包含该键,加载时不会报错。
方法3:宽松加载(不推荐)
加载时设置strict=False忽略所有不匹配的键,但此方法可能隐藏其他参数匹配问题,仅用于临时测试:
# 使用load_from_checkpoint时 model = QTagClassifier.load_from_checkpoint("example.ckpt", strict=False) # 使用trainer.fit时 trainer.fit(model, QTdata_module, ckpt_path="example.ckpt", strict=False)
内容的提问来源于stack exchange,提问作者Ali Nazeri
相关产品推荐
相关产品推荐

