为ONNX BERT模型添加Reshape节点时如何设置动态维度?
问题描述
尝试为支持动态形状的BERT ONNX模型(opset 18)添加Reshape节点,目标是将形状为[unk__2,unk__3,768]的3阶张量重塑为2阶张量,合并前两个动态维度为[unk__2 * unk__3],保留最后一个固定维度768。使用ONNX Helper创建张量时触发报错,代码片段如下:
# Create a Constant node that contains the target shape shape_tensor = helper.make_tensor_value_info(name='shape', elem_type=onnx.TensorProto.INT64, shape=(-1,768)) shape_node = helper.make_node( 'Constant', inputs=[], outputs=[f'shape_{i}_output'], value=shape_tensor, name=f'shape_{i}' ) # Create a Reshape node reshape_node = helper.make_node( 'Reshape', inputs=[mm_node.input[0], f'shape_{i}_output'], outputs=[f'reshaped_output_{i}'], name=f'Reshape_{i}' )
运行时错误:
raise TypeError(f"'{value}' is not an accepted attribute value.") TypeError: 'name: "shape" type { tensor_type { elem_type: 7 shape { dim { dim_value: -1 } dim { dim_value: 768 } } } } ' is not an accepted attribute value.
解决方法
报错核心原因:helper.make_tensor_value_info()生成的是张量类型描述信息,但Constant节点的value参数需要的是包含具体数值的TensorProto实例。针对动态维度合并需求,提供两种可行方案:
方法1:用Reshape的自动推导维度(推荐)
ONNX opset 18的Reshape算子支持用-1自动推导合并后的维度,只需构造包含[-1, 768]的常量张量即可:
import numpy as np # 构造包含目标形状的TensorProto实例 shape_data = np.array([-1, 768], dtype=np.int64) shape_tensor = helper.make_tensor( name=f'shape_tensor_{i}', data_type=onnx.TensorProto.INT64, dims=shape_data.shape, vals=shape_data.flatten().tolist() ) # 创建Constant节点 shape_node = helper.make_node( 'Constant', inputs=[], outputs=[f'shape_{i}_output'], value=shape_tensor, name=f'shape_{i}' ) # 创建Reshape节点(逻辑不变) reshape_node = helper.make_node( 'Reshape', inputs=[mm_node.input[0], f'shape_{i}_output'], outputs=[f'reshaped_output_{i}'], name=f'Reshape_{i}' )
方法2:显式计算合并维度(适合需明确维度值的场景)
如果需要显式计算unk__2 * unk__3的结果,可通过Shape、Gather、Mul算子组合实现:
# 1. 获取输入张量的形状 shape_input_node = helper.make_node( 'Shape', inputs=[mm_node.input[0]], outputs=[f'input_shape_{i}'], name=f'Shape_{i}' ) # 2. 提取前两个动态维度 dim0_node = helper.make_node( 'Gather', inputs=[f'input_shape_{i}', helper.make_tensor('axis0', onnx.TensorProto.INT64, [1], [0])], outputs=[f'dim0_{i}'], name=f'GatherDim0_{i}' ) dim1_node = helper.make_node( 'Gather', inputs=[f'input_shape_{i}', helper.make_tensor('axis1', onnx.TensorProto.INT64, [1], [1])], outputs=[f'dim1_{i}'], name=f'GatherDim1_{i}' ) # 3. 计算两个维度的乘积 merged_dim_node = helper.make_node( 'Mul', inputs=[f'dim0_{i}', f'dim1_{i}'], outputs=[f'merged_dim_{i}'], name=f'MulDims_{i}' ) # 4. 拼接得到目标形状 target_shape_node = helper.make_node( 'Concat', inputs=[f'merged_dim_{i}', helper.make_tensor('fixed_dim', onnx.TensorProto.INT64, [1], [768])], outputs=[f'target_shape_{i}'], axis=0, name=f'ConcatShape_{i}' ) # 5. 创建Reshape节点 reshape_node = helper.make_node( 'Reshape', inputs=[mm_node.input[0], f'target_shape_{i}'], outputs=[f'reshaped_output_{i}'], name=f'Reshape_{i}' )
关键注意事项
helper.make_tensor_value_info()仅用于定义张量的输入/输出类型,不能作为Constant节点的value参数;helper.make_tensor()才是生成包含具体数值的TensorProto,符合Constant节点的参数要求;- opset 18完全兼容Reshape的
-1自动推导逻辑,方法1更简洁高效。
内容的提问来源于stack exchange,提问作者linos
相关产品推荐
相关产品推荐

