模型成功转ONNX后Onnxruntime测试报错问题咨询
问题:PyTorch Transformer模型转ONNX后动态输入运行报错
问题重现
模型基于PyTorch的nn.Transformer构建,仅调用其encoder部分。导出ONNX时使用固定形状的输入张量((12,1,100)),指定了输入第0轴为动态轴,但运行时输入其他长度的张量(如(10,1,100))会触发Reshape错误,只有导出时的输入形状能正常运行。
完整代码:
#!/usr/bin/env python3 import torch.nn as nn from torch import Tensor, rand import torch.onnx import onnx onnx_model = 'MicroTest.onnx' class Test_trans(nn.Module): def __init__( self, emb_size=100): super(Test_trans, self).__init__() self.transformer = nn.Transformer(emb_size, 2, 2, 2, 512, 0.1) def forward(self, src: Tensor): return self.transformer.encoder(src) def process_one_torch(session, ten): print('Tensor In size:', ten.size(), end='\t') memory = session(ten) print('Mem size:', memory.size()); def process_one_onnx(session, npa): ortvalue = onnxruntime.OrtValue.ortvalue_from_numpy(npa) print('In ortvalue.shape:', ortvalue.shape(), end='\t') memory = session.run(None, {session.get_inputs()[0].name: ortvalue}) print('ONNX mem.shape:', memory[0].shape) mini = Test_trans() c_tensor_12 = rand((12,1,100)) c_tensor_10 = rand((10,1,100)) print('################################# Torch ###################################') process_one_torch(mini, c_tensor_12) process_one_torch(mini, c_tensor_10) torch.onnx.export(mini, # model being run c_tensor_12, # model input (or a tuple for multiple inputs) onnx_model, # where to save the model (can be a file or file-like object) export_params=True, # store the trained parameter weights inside the model file opset_version=17, # the ONNX version to export the model to do_constant_folding=False, input_names = ['input'], # the model's input names output_names = ['output'], # the model's output names dynamic_axes = {'input' : {0: 'max_len'}}) print('################################# ONNX RT #################################') import onnxruntime session = onnxruntime.InferenceSession(onnx_model, providers=["CPUExecutionProvider"]) print('Session inputs:', session.get_inputs()[0]) process_one_onnx(session, c_tensor_12.numpy()) process_one_onnx(session, c_tensor_10.numpy()) #This one crashes
报错详情
onnxruntime_test MicroTest.onnx 2024-06-12 17:11:22.877402346 [E:onnxruntime:, sequential_executor.cc:514 ExecuteKernel] Non-zero status code returned while running Reshape node. Name:'/encoder/layers.0/self_attn/Reshape_4' Status Message: /croot/onnxruntime_1711063034809/work/onnxruntime/core/providers/cpu/tensor/reshape_helper.h:44 onnxruntime::ReshapeHelper::ReshapeHelper(const onnxruntime::TensorShape&, onnxruntime::TensorShapeVector&, bool) input_shape_size == size was false. The input tensor cannot be reshaped to the requested shape. Input shape:{1,1,100}, requested shape:{12,2,50} Traceback (most recent call last): File "/home/if/miniconda3/envs/cpu/bin/onnxruntime_test", line 11, in sys.exit(main()) ^^^^^^ File "/home/if/miniconda3/envs/cpu/lib/python3.11/site-packages/onnxruntime/tools/onnxruntime_test.py", line 159, in main exit_code, _, _ = run_model(args.model_path, args.num_iters, args.debug, args.profile, args.symbolic_dims) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ File "/home/if/miniconda3/envs/cpu/lib/python3.11/site-packages/onnxruntime/tools/onnxruntime_test.py", line 118, in run_model outputs = sess.run([], feeds) # fetch all outputs ^^^^^^^^^^^^^^^^^^^ File "/home/if/miniconda3/envs/cpu/lib/python3.11/site-packages/onnxruntime/capi/onnxruntime_inference_collection.py", line 220, in run return self._sess.run(output_names, input_feed, run_options) ^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ onnxruntime.capi.onnxruntime_pybind11_state.RuntimeException: [ONNXRuntimeError] : 6 : RUNTIME_EXCEPTION : Non-zero status code returned while running Reshape node. Name:'/encoder/layers.0/self_attn/Reshape_4' Status Message: /croot/onnxruntime_1711063034809/work/onnxruntime/core/providers/cpu/tensor/reshape_helper.h:44 onnxruntime::ReshapeHelper::ReshapeHelper(const onnxruntime::TensorShape&, onnxruntime::TensorShapeVector&, bool) input_shape_size == size was false. The input tensor cannot be reshaped to the requested shape. Input shape:{1,1,100}, requested shape:{12,2,50}
原因分析
- 动态轴未完全覆盖:仅标记输入的第0轴为动态,但Transformer内部多头注意力模块的Reshape操作被PyTorch导出时硬编码为导出示例的输入长度(12),未关联到动态维度符号。当输入长度变为10时,Reshape目标形状仍为
(12,2,50),与输入张量的元素总数(11100=100)不匹配(12250=1200),触发错误。 - Transformer内部结构的导出限制:PyTorch的
nn.Transformer内部实现中,部分形状推导依赖静态输入形状,导出ONNX时若仅指定输入动态轴,PyTorch无法自动将所有依赖该维度的内部Reshape操作转换为动态形状。
解决方案
1. 完善动态轴配置
同时标记输入和输出的动态轴,确保ONNX追踪所有关联维度:
torch.onnx.export(mini, c_tensor_12, onnx_model, export_params=True, opset_version=17, do_constant_folding=False, input_names=['input'], output_names=['output'], # 同时标记输入和输出的第0轴为动态 dynamic_axes={ 'input': {0: 'max_len'}, 'output': {0: 'max_len'} })
2. 使用多形状示例输入导出
PyTorch支持提供多个示例输入,帮助导出器识别动态维度的变化范围,避免硬编码:
# 准备多个不同长度的示例输入 dummy_inputs = [rand((12,1,100)), rand((10,1,100))] torch.onnx.export(mini, dummy_inputs, onnx_model, export_params=True, opset_version=17, do_constant_folding=False, input_names=['input'], output_names=['output'], dynamic_axes={ 'input': {0: 'max_len'}, 'output': {0: 'max_len'} })
3. 自定义Transformer Encoder实现
若上述方法无效,可手动实现Transformer Encoder的核心逻辑,显式用动态形状计算Reshape目标,避免硬编码:
class CustomTransformerEncoder(nn.Module): def __init__(self, emb_size=100, num_heads=2, num_layers=2): super().__init__() self.layer = nn.TransformerEncoderLayer(emb_size, num_heads, dim_feedforward=512, dropout=0.1) self.encoder = nn.TransformerEncoder(self.layer, num_layers) def forward(self, src: Tensor): seq_len, batch_size, emb_size = src.size() # 显式确保所有内部操作基于输入的动态形状 return self.encoder(src)
4. 验证ONNX模型动态性
用Netron工具打开导出的ONNX模型,检查Reshape节点的目标形状是否为动态符号(如max_len)而非固定值12,确认动态维度已正确标记。
内容的提问来源于stack exchange,提问作者Dodiak
相关产品推荐
相关产品推荐

