在两个ONNX模型上添加If节点时出现形状推断错误求排查
问题分析与解决方案
错误原因
你遇到的错误是因为ONNX If 节点的分支子图(then_branch/else_branch)没有正确接收主图传递的输入。具体来说:
If节点的规范要求:分支子图的输入数量必须等于If节点的输入列表中除条件张量之外的输入数量- 你的代码中
If节点仅传入了条件张量"split",但分支子图(split和nosplit模型的图)本身需要接收图像输入,导致子图期望1个输入但实际未收到任何输入,触发形状推断错误。
修正后的代码
import onnx # 加载原模型并校验 onnx_model_split = onnx.load("split.onnx") onnx_model_no_split = onnx.load("nosplit.onnx") onnx.checker.check_model(onnx_model_no_split, full_check=True) onnx.checker.check_model(onnx_model_split, full_check=True) assert len(onnx_model_split.graph.output) == len(onnx_model_no_split.graph.output) # 获取原模型的输入(图像等非条件输入) original_inputs = list(onnx_model_no_split.graph.input) # 创建分支图:保留原模型的节点、输入、输出、初始化器和中间张量信息 graph_split: onnx.GraphProto = onnx.helper.make_graph( nodes=list(onnx_model_split.graph.node), name="graph-split", inputs=list(onnx_model_split.graph.input), outputs=list(onnx_model_split.graph.output), initializer=list(onnx_model_split.graph.initializer), value_info=list(onnx_model_split.graph.value_info), # 保留中间张量的形状/类型信息 ) graph_no_split: onnx.GraphProto = onnx.helper.make_graph( nodes=list(onnx_model_no_split.graph.node), name="graph-no-split", inputs=list(onnx_model_no_split.graph.input), outputs=list(onnx_model_no_split.graph.output), initializer=list(onnx_model_no_split.graph.initializer), value_info=list(onnx_model_no_split.graph.value_info), # 保留中间张量的形状/类型信息 ) # 定义条件输入的张量信息 split_input = onnx.helper.make_tensor_value_info(name="split", elem_type=onnx.TensorProto.BOOL, shape=[1]) # 创建If节点:传入条件张量+原模型输入,确保分支能收到所需输入 if_node = onnx.helper.make_node( op_type="If", inputs=["split"] + [inp.name for inp in original_inputs], # 加入原模型的输入 outputs=[o.name for o in onnx_model_no_split.graph.output], then_branch=graph_split, else_branch=graph_no_split, ) # 创建主图:输入包含原模型输入+条件输入,初始化器留空(分支图自带初始化器) if_graph_def: onnx.GraphProto = onnx.helper.make_graph( nodes=[if_node], name="if-model", inputs=original_inputs + [split_input], outputs=list(onnx_model_no_split.graph.output), initializer=[], # 主图不需要重复包含分支的初始化器 ) # 生成最终模型,指定opset 12(If节点的要求) model_def: onnx.ModelProto = onnx.helper.make_model( if_graph_def, producer_name="onnx-example", opset_imports=[onnx.helper.make_opsetid(onnx.defs.ONNX_DOMAIN, 12)] ) # 校验并保存模型 onnx.checker.check_model(model_def, full_check=True) onnx.save(model_def, 'test.onnx')
关键修改点
- If节点输入补充:将原模型的图像输入加入
If节点的inputs列表,确保分支子图能接收到所需的输入张量 - 保留分支图的value_info:原模型的中间张量形状/类型信息能帮助ONNX进行准确的形状推断,避免潜在错误
- 主图初始化器去重:分支图已自带原模型的初始化器,主图无需重复添加,避免张量名称冲突
内容的提问来源于stack exchange,提问作者ahmed belgacem
相关产品推荐
相关产品推荐

