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

将PyTorch版GANet导出为ONNX时遇RuntimeError问题求助

问题

在Python 3.9.13 + PyTorch 1.13.0环境下,将PyTorch实现的GANet模型导出为ONNX文件时,触发RuntimeError: Argument passed to at() was not in the map错误。

导出代码如下:

# 训练设置
parser = argparse.ArgumentParser(description='PyTorch GANet Example')
parser.add_argument('--crop_height', type=int, required=True, help="crop height")
parser.add_argument('--crop_width', type=int, required=True, help="crop width")
parser.add_argument('--max_disp', type=int, default=192, help="max disp")
parser.add_argument('--resume', type=str, default='', help="resume from saved model")
parser.add_argument('--cuda', type=bool, default=True, help='use cuda?')
parser.add_argument('--kitti', type=int, default=0, help='kitti dataset? Default=False')
parser.add_argument('--kitti2015', type=int, default=0, help='kitti 2015? Default=False')
parser.add_argument('--data_path', type=str, required=True, help="data root")
parser.add_argument('--test_list', type=str, required=True, help="training list")
parser.add_argument('--save_path', type=str, default='./result/', help="location to save result")
parser.add_argument('--model', type=str, default='GANet_deep', help="model to train")

opt = parser.parse_args()


model = GANet(opt.max_disp)

print("=> loading checkpoint '{}'".format(opt.resume))
checkpoint = torch.load(opt.resume)
model.load_state_dict(checkpoint['state_dict'], strict=False)


model.eval()

# 模型输入
x = torch.randn(1, 3, 48, 48, requires_grad=True)
torch_out = model(x,x)
batch_size = 1


# 导出模型
torch.onnx.export(model,               # 待导出模型
                  (x,x),                         # 模型输入(多输入用元组)
                  "onnxGanet.onnx",   # 模型保存路径
                  export_params=True,        # 保存训练好的参数
                  opset_version=11,          # ONNX版本
                  do_constant_folding=True,  # 执行常量折叠优化
                  input_names = ['input'],   # 输入名称
                  output_names = ['output'], # 输出名称
                  dynamic_axes={'input' : {0 : 'batch_size'},    # 动态轴
                                'output' : {0 : 'batch_size'}})
                                 
import onnx

onnx_model = onnx.load("onnxGanet.onnx")

# 保存ONNX模型
onnx.save(onnx_model, "/home/jokar/GANet-master/onnxGanet.onnx")

报错信息:

/home/jokar/GANet-master/onnxGANet.py:95: TracerWarning: Converting a tensor to a Python boolean might cause the trace to be incorrect. We can't record the data flow of Python values, so this value will be treated as a constant in the future. This means that the trace might not generalize to other inputs!
  assert(x.size() == rem.size())
/home/jokar/GANet-master/libs/GANet/modules/GANet.py:123: TracerWarning: resize_ can't be represented in the JIT at the moment, so we won't connect any uses of this value with its current trace. If you happen to use it again, it will show up as a constant in the graph. Consider using `view` or `reshape` to make it traceable.
  cost = x.new().resize_(num, channels * 2, self.maxdisp, height, width).zero_()
/home/jokar/GANet-master/libs/GANet/functions/GANet.py:14: TracerWarning: resize_ can't be represented in the JIT at the moment, so we won't connect any uses of this value with its current trace. If you happen to use it again, it will show up as a constant in the graph. Consider using `view` or `reshape` to make it traceable.
  output = input.new().resize_(num, channels, depth, height, width).zero_()
/home/jokar/GANet-master/libs/GANet/functions/GANet.py:15: TracerWarning: resize_ can't be represented in the JIT at the moment, so we won't connect any uses of this value with its current trace. If you happen to use it again, it will show up as a constant in the graph. Consider using `view` or `reshape` to make it traceable.
  temp_out = input.new().resize_(num, channels, depth, height, width).zero_()
/home/jokar/GANet-master/libs/GANet/functions/GANet.py:16: TracerWarning: resize_ can't be represented in the JIT at the moment, so we won't connect any uses of this value with its current trace. If you happen to use it again, it will show up as a constant in the graph. Consider using `view` or `reshape` to make it traceable.
  mask = input.new().resize_(num, channels, depth, height, width).zero_()
/home/jokar/GANet-master/onnxGANet.py:305: TracerWarning: Converting a tensor to a Python boolean might cause the trace to be incorrect. We can't record the data flow of Python values, so this value will be treated as a constant in the future. This means that the trace might not generalize to other inputs!
  assert(x.size() == rem.size())
/home/jokar/GANet-master/onnxGANet.py:272: TracerWarning: Converting a tensor to a Python boolean might cause the trace to be incorrect. We can't record the data flow of Python values, so this value will be treated as a constant in the future. This means that the trace might not generalize to other inputs!
  assert(lg1.size() == lg2.size())
/home/jokar/GANet-master/libs/GANet/functions/GANet.py:181: TracerWarning: resize_ can't be represented in the JIT at the moment, so we won't connect any uses of this value with its current trace. If you happen to use it again, it will show up as a constant in the graph. Consider using `view` or `reshape` to make it traceable.
  temp_out = input.new().resize_(num, channels, height, width).zero_()
/home/jokar/GANet-master/libs/GANet/functions/GANet.py:182: TracerWarning: resize_ can't be represented in the JIT at the moment, so we won't connect any uses of this value with its current trace. If you happen to use it again, it will show up as a constant in the graph. Consider using `view` or `reshape` to make it traceable.
  output = input.new().resize_(num, channels, height, width).zero_()
/home/jokar/GANet-master/libs/GANet/modules/GANet.py:145: TracerWarning: torch.Tensor results are registered as constants in the trace. You can safely ignore this warning if you use this function to create tensors out of constant variables that would be the same every time you call this function. In any other case, this might cause the trace to be incorrect.
  disp = Variable(torch.Tensor(np.reshape(np.array(range(self.maxdisp)),[1, self.maxdisp, 1, 1])), requires_grad=False)
Traceback (most recent call last):
  File "/home/jokar/GANet-master/onnxGANet.py", line 481, in <module>
    torch.onnx.export(model,               # model being run
  File "/opt/anaconda3/lib/python3.9/site-packages/torch/onnx/utils.py", line 504, in export
    _export(
  File "/opt/anaconda3/lib/python3.9/site-packages/torch/onnx/utils.py", line 1529, in _export
    graph, params_dict, torch_out = _model_to_graph(
  File "/opt/anaconda3/lib/python3.9/site-packages/torch/onnx/utils.py", line 1115, in _model_to_graph
    graph = _optimize_graph(
  File "/opt/anaconda3/lib/python3.9/site-packages/torch/onnx/utils.py", line 617, in _optimize_graph
    _C._jit_pass_onnx_remove_inplace_ops_for_onnx(graph, module)
RuntimeError: Argument passed to at() was not in the map.

解决方案

该错误源于模型中不可追踪的原地操作和Python层面的张量判断,导致PyTorch JIT无法正确生成ONNX图,按以下步骤修复:

1. 替换所有resize_原地操作

resize_是JIT无法追踪的原地操作,改用torch.zeros直接创建指定形状的张量,替代new().resize_().zero_()的写法:

原代码:

cost = x.new().resize_(num, channels * 2, self.maxdisp, height, width).zero_()

修改为:

cost = torch.zeros((num, channels * 2, self.maxdisp, height, width), device=x.device, dtype=x.dtype)

对所有出现resize_的文件(libs/GANet/modules/GANet.py、libs/GANet/functions/GANet.py等)做同样替换。

2. 替换Python断言为张量断言

原代码中的assert(x.size() == rem.size())会将张量尺寸转为Python布尔值,破坏追踪流程,改用张量比较:

原代码:

assert(x.size() == rem.size())

修改为:

assert torch.all(torch.tensor(x.size()) == torch.tensor(rem.size())), "Size mismatch between tensors"

3. 修正常量张量的创建方式

原代码创建disp的方式会被注册为常量,影响动态输入兼容性,改用torch.arange直接生成:

原代码:

disp = Variable(torch.Tensor(np.reshape(np.array(range(self.maxdisp)),[1, self.maxdisp, 1, 1])), requires_grad=False)

修改为:

disp = torch.arange(self.maxdisp, device=x.device).view(1, self.maxdisp, 1, 1).detach()

(注:PyTorch 1.0+已弃用Variable,直接使用张量即可)

4. 调整ONNX导出参数

  • 提升opset_version到13(PyTorch 1.13支持opset 13+,对原地操作兼容性更好)
  • 暂时关闭do_constant_folding,避免折叠过程出错
  • 区分两个输入的名称(原代码中两个输入都叫input会导致冲突)

修改后的导出代码片段:

# 导出模型
torch.onnx.export(model,               
                  (x,x),                         
                  "onnxGanet.onnx",   
                  export_params=True,        
                  opset_version=13,          
                  do_constant_folding=False,  
                  input_names = ['left_input', 'right_input'],
                  output_names = ['output'], 
                  dynamic_axes={'left_input' : {0 : 'batch_size'},    
                                'right_input' : {0 : 'batch_size'},
                                'output' : {0 : 'batch_size'}})

5. 验证导出的ONNX模型

导出完成后,用ONNX工具验证模型合法性并测试推理:

import onnx
from onnxruntime import InferenceSession
import torch

# 加载并检查模型
onnx_model = onnx.load("onnxGanet.onnx")
onnx.checker.check_model(onnx_model)

# 用ONNX Runtime测试推理
session = InferenceSession("onnxGanet.onnx")
input_left = x.numpy()
input_right = x.numpy()
ort_output = session.run(None, {"left_input": input_left, "right_input": input_right})

# 对比PyTorch和ONNX Runtime的输出
torch.testing.assert_close(torch_out[0].detach().numpy(), ort_output[0], rtol=1e-3, atol=1e-3)

内容的提问来源于stack exchange,提问作者user297962

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.08.09 22:10:31