PyTorch预训练.pth模型转TensorFlow Protobuf遇ONNX转换错误求助
搞定PyTorch→ONNX→TensorFlow Protobuf的模型转换问题
嘿,我来帮你解决这个模型格式转换的难题!你在PyTorch转ONNX环节出错,大概率是模型加载或者导出代码的细节没处理到位,我给你梳理一套完整的排查和解决流程,一步步来:
第一步:先把PyTorch模型加载的坑填上
你用model = torch.load('./my_model.pth')加载模型时,得先搞清楚你的.pth文件是权重文件还是完整模型文件:
- 如果是权重文件(只有参数,没有模型结构):你必须先定义好和训练时完全一致的模型结构,再加载权重,而且一定要切换到评估模式(避免BatchNorm、Dropout这类层搞事情):
# 举个例子,假设你的模型是自定义的ResNet50,这里要写你自己的模型结构定义 model = torchvision.models.resnet50(pretrained=False) # 加载权重,注意如果权重是多GPU训练的,可能需要加map_location='cpu'或者去掉module前缀 model.load_state_dict(torch.load('./my_model.pth', map_location='cpu')) # 关键!切换到评估模式 model.eval() - 如果是完整模型文件:加载后同样要切到评估模式:
model = torch.load('./my_model.pth', map_location='cpu') model.eval()
第二步:修正ONNX导出的代码
你的原代码有几个过时或者不严谨的地方,我给你调整成适配新版本PyTorch的写法:
import torch import torch.onnx # 构造dummy输入,shape要和模型实际输入完全匹配(这里是[1,3,256,256],你根据自己的模型改) dummy_input = torch.randn(1, 3, 256, 256) # 上面已经加载好并设置为eval模式的model # model = ...(参考第一步的代码) # 导出ONNX模型,参数都给你标清楚了 torch.onnx.export( model, dummy_input, "my_model.onnx", # 输出的ONNX文件名 export_params=True, # 导出带权重的完整模型 opset_version=13, # 选个较高的opset版本,兼容性更好,比如13或11 do_constant_folding=True, # 开启常量折叠优化 input_names=['input'], # 给输入节点起个名字,后续转TensorFlow方便 output_names=['output'], # 给输出节点起名字 dynamic_axes={'input': {0: 'batch_size'}, 'output': {0: 'batch_size'}} # 支持动态batch大小,可选但实用 )
第三步:验证ONNX模型是否正常
导出后别急着转TensorFlow,先验证下ONNX模型能不能跑:
import onnxruntime as ort import numpy as np # 加载ONNX模型 ort_session = ort.InferenceSession("my_model.onnx") # 造个测试输入,和dummy_input的shape、 dtype一致 test_input = np.random.randn(1, 3, 256, 256).astype(np.float32) # 跑推理 outputs = ort_session.run(None, {'input': test_input}) print("输出shape:", outputs[0].shape) # 看看输出是不是你预期的样子
第四步:把ONNX转成TensorFlow Protobuf
确认ONNX模型没问题后,用onnx-tf工具转TensorFlow:
- 先装依赖:
pip install onnx-tf - 命令行执行转换,生成TensorFlow的SavedModel:
onnx-tf convert -i my_model.onnx -o my_model_tf_dir - 如果需要转成传统的.pb冻结图格式,再用这段代码处理:
import tensorflow as tf from tensorflow.python.framework.convert_to_constants import convert_variables_to_constants_v2 # 加载SavedModel saved_model = tf.saved_model.load('my_model_tf_dir') infer_func = saved_model.signatures['serving_default'] # 把变量都转成常量,冻结图 frozen_func = convert_variables_to_constants_v2(infer_func) frozen_graph = frozen_func.graph # 保存成.pb文件 tf.io.write_graph( graph_or_graph_def=frozen_graph, logdir='.', name='my_model_frozen.pb', as_text=False )
常见错误排查小技巧
- 导出ONNX时提示算子不支持:试试调整
opset_version,如果是自定义算子,得额外写ONNX的算子导出逻辑 - 加载权重时key不匹配:说明你定义的模型结构和训练时不一样,要确保每一层的名称、结构完全一致,多GPU训练的权重可能需要去掉
module.前缀 - 转TensorFlow时出错:检查onnx和onnx-tf的版本是否兼容,建议更到最新版试试
内容的提问来源于stack exchange,提问作者Urvish
相关产品推荐
相关产品推荐

