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

模型成功转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}

原因分析

  1. 动态轴未完全覆盖:仅标记输入的第0轴为动态,但Transformer内部多头注意力模块的Reshape操作被PyTorch导出时硬编码为导出示例的输入长度(12),未关联到动态维度符号。当输入长度变为10时,Reshape目标形状仍为(12,2,50),与输入张量的元素总数(11100=100)不匹配(12250=1200),触发错误。
  2. 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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.22 12:25:55