如何将PyTorch版D3NET转为ONNX?转换代码无输出无报错求助
PyTorch转ONNX无输出无报错问题排查与解决
问题描述
近期使用D3NetBenchmark模型训练深度图像,得到三个最优.pth模型文件。参照PyTorch官方教程编写了Export_onnx.py代码尝试转换其中一个模型,代码如下:
import torch.onnx from model.RgbNet import MyNet as RgbNet pretrained_models={'RgbNet':'./eval/pretrained_models/RgbNet.pth', 'RgbdNet':'eval/pretrained_models/RgbdNet.pth' , 'DepthNet':'eval/pretrained_models/DepthNet.pth' } model_rgbd=RgbdNet() model_rgb.load_state_dict(torch.load(pretrained_models['RgbNet'])['model']) model_rgb.eval() #Dummy Input 1 = RGB Image, Dummy Input 2 = Depth Image batch_size = 1 dummy_input1 = torch.randn(batch_size, 3, 224, 224, requires_grad=True, dtype=torch.float32) dummy_input2 = torch.randn(batch_size, 1, 224, 224, requires_grad=True, dtype=torch.float32) input = (dummy_input1, dummy_input2) torch.onnx.export(model_rgbd.cpu(), (input,), "Model.onnx")
运行代码后无任何输出,也未报错,仅终端显示网络扫描相关信息,需要解决该问题。
解决方案
1. 修正模型实例与权重加载的匹配问题
代码中存在模型变量定义混乱的问题:
- 实例化了
model_rgbd但未加载权重,同时引用了未定义的model_rgb进行权重加载 - 若目标是转换
RgbNet,修正代码如下:# 正确导入并实例化RgbNet model_rgb = RgbNet() # 加载对应权重 model_rgb.load_state_dict(torch.load(pretrained_models['RgbNet'])['model']) model_rgb.eval() - 若目标是转换
RgbdNet,需先导入该模型类,再对应加载RgbdNet.pth的权重。
2. 修正ONNX导出的输入参数格式
torch.onnx.export的输入参数被错误嵌套了一层元组,导致模型无法正确接收输入:
- 原代码中
(input,)会将输入包装成((dummy_input1, dummy_input2),),与模型预期的双输入不匹配 - 修正为直接传入
input:torch.onnx.export(model_rgb.cpu(), input, "Model.onnx")
3. 添加导出日志与模型验证
- 开启
verbose=True参数,打印导出过程的详细日志,便于排查潜在问题:torch.onnx.export(model_rgb.cpu(), input, "Model.onnx", verbose=True) - 导出完成后,用ONNX Runtime验证模型有效性:
import onnxruntime as ort import numpy as np # 加载模型 sess = ort.InferenceSession("Model.onnx") input_names = [node.name for node in sess.get_inputs()] output_names = [node.name for node in sess.get_outputs()] print("模型输入节点:", input_names) print("模型输出节点:", output_names) # 用dummy input测试推理 test_output = sess.run(output_names, { input_names[0]: dummy_input1.numpy(), input_names[1]: dummy_input2.numpy() }) print("推理输出形状:", [out.shape for out in test_output])
4. 确认模型导入路径正确性
确保model.RgbNet或model.RgbdNet的导入路径符合当前运行环境的目录结构,避免因找不到模型类导致的隐性错误。
内容的提问来源于stack exchange,提问作者Steven Yen
相关产品推荐
相关产品推荐

