加载PyTorch训练的model_dict.pt文件时出现TypeError错误求助
解决PyTorch模型加载TypeError问题
问题根源
你保存模型时直接存储了整个BertForTokenClassification模型对象,而非模型的state_dict(参数字典)。torch.load()读取后返回的是完整模型实例,但load_state_dict()方法要求传入的是键值对形式的参数字典,因此触发类型错误。
两种解决方法
方法1:调整加载代码(无需重新保存)
如果无法重新生成保存文件,直接加载整个模型对象即可,不用调用load_state_dict():model = torch.load(path+'model_dict.pt', map_location=device)注意:这种方式对环境兼容性要求高,需要加载环境和保存环境的PyTorch、transformers版本一致。
方法2:重新保存模型(推荐)
若能重新执行保存步骤,修改保存代码为存储模型的state_dict:# 保存模型参数(替换原来的保存代码) torch.save(model.state_dict(), path+'model_dict.pt') # 加载时使用你原本的代码即可正常运行 model.load_state_dict(torch.load(path+'model_dict.pt', map_location=device))
补充说明
state_dict是PyTorch官方推荐的模型参数存储格式,仅包含可学习参数和缓冲区数据,文件体积更小,跨环境兼容性更强。- 直接保存模型对象会包含模型结构、环境依赖等额外信息,迁移到不同版本的框架环境时容易出现加载失败。
内容的提问来源于stack exchange,提问作者박찬우
相关产品推荐
相关产品推荐

