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

在两个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')

关键修改点

  1. If节点输入补充:将原模型的图像输入加入If节点的inputs列表,确保分支子图能接收到所需的输入张量
  2. 保留分支图的value_info:原模型的中间张量形状/类型信息能帮助ONNX进行准确的形状推断,避免潜在错误
  3. 主图初始化器去重:分支图已自带原模型的初始化器,主图无需重复添加,避免张量名称冲突

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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.20 16:24:54