You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

如何将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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.28 18:50:36