Pytorch手写文本识别预训练编码解码模型预测TypeError问题咨询
问题根因
你触发报错的核心原因是混淆了模型权重字典(state_dict)和模型实例的概念:checkpoint['encoder_state_dict']和checkpoint['decoder_state_dict']存储的是模型的权重参数,是OrderedDict类型的静态数据,不是可执行前向推理的模型对象,自然不能被当作函数调用。
解决步骤
- 第一步:先实例化你用的SPAN模型对应的编码器、解码器类,注意类的结构、参数要和预训练模型训练时的配置完全一致,否则会出现权重加载不匹配的问题
- 第二步:把加载到的state_dict权重加载到实例化后的模型对象上,再执行推理
- 修正后的代码示例如下:
import torch # 首先导入你项目里的Encoder、Decoder类,要和预训练用的类定义完全一致 from your_model_file import Encoder, Decoder # 1. 加载预训练权重 checkpoint = torch.load("Model/SPAN/SPAN-PT-RA_rimes.pt",map_location=torch.device('cpu')) encoder_state_dict = checkpoint['encoder_state_dict'] decoder_state_dict = checkpoint['decoder_state_dict'] # 2. 实例化模型,构造参数按训练时的配置填写,比如输入通道、层数等 encoder = Encoder() decoder = Decoder() # 3. 加载权重到模型实例 encoder.load_state_dict(encoder_state_dict) decoder.load_state_dict(decoder_state_dict) # 4. 切换到评估模式,避免dropout、BN层推理时和训练行为不一致 encoder.eval() decoder.eval() # 5. 执行推理,关闭梯度计算减少资源占用 img = torch.LongTensor(img).unsqueeze(1).to(torch.device('cpu')) with torch.no_grad(): encoder_out = encoder(img) global_pred = decoder(encoder_out)
- 额外注意:如果加载权重时出现key不匹配的警告/报错,先检查你实例化模型的参数是否和预训练训练时的配置完全一致;如果是分布式训练存储的权重key多了
module.前缀的问题,手动去掉key前缀再加载即可。
内容的提问来源于stack exchange,提问作者Imen Bouzidi
相关产品推荐
相关产品推荐

