将PyTorch转CoreML遇TypeError,无模型类时如何解决?
解决本地PyTorch模型加载报错及无模型类时的转换方案
错误原因
torch.load('local_model_file.pth')加载的是模型的状态字典(state_dict),这是一个存储模型参数的字典对象,并非可直接调用的模型实例,因此执行model(input_batch)会抛出dict object is not callable错误。
无模型类时的解决办法
方法1:通过Torch Hub加载模型结构再导入本地参数
如果你的模型来自公开仓库(比如PyTorch Vision),可以先从Torch Hub加载对应模型的结构(不加载预训练权重),再将本地state_dict导入:
# 加载模型结构(pretrained=False表示不加载官方预训练权重) model = torch.hub.load('pytorch/vision:v0.6.0', 'deeplabv3_resnet101', pretrained=False).eval() # 加载本地保存的state_dict state_dict = torch.load('local_model_file.pth') # 将参数加载到模型中 model.load_state_dict(state_dict) # 后续推理流程不变 input_tensor = preprocess(input_image) input_batch = input_tensor.unsqueeze(0) with torch.no_grad(): output = model(input_batch)['out'][0] torch_predictions = output.argmax(0)
方法2:尝试加载TorchScript格式模型
如果本地模型是通过torch.jit.save导出的TorchScript格式,可直接用torch.jit.load加载为可调用的模型实例:
model = torch.jit.load('local_model_file.pth').eval() # 推理代码保持一致 input_tensor = preprocess(input_image) input_batch = input_tensor.unsqueeze(0) with torch.no_grad(): output = model(input_batch)['out'][0] torch_predictions = output.argmax(0)
方法3:从state_dict反推模型结构
若上述方法都不可行,可通过分析state_dict的参数键名推断模型结构:
state_dict = torch.load('local_model_file.pth') # 打印部分参数键名,观察特征(比如是否包含backbone、classifier等字段) print(list(state_dict.keys())[:10])
根据键名特征匹配公开的模型架构(比如键名含deeplabv3则对应DeepLabV3系列),创建对应模型类实例后再加载state_dict。
内容的提问来源于stack exchange,提问作者AnupamChugh
相关产品推荐
相关产品推荐

