You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.11 21:28:20