This is an automated email from the ASF dual-hosted git repository.
tlopex pushed a commit to branch main
in repository https://gitbox.apache.org/repos/asf/tvm.git
The following commit(s) were added to refs/heads/main by this push:
new 0e75b43a62 [Fix][Relax][ONNX] Relax op normalization for onnx
subgraphs (#20010)
0e75b43a62 is described below
commit 0e75b43a62c31d0f14aa6619462fa845d2c5b5a3
Author: Balint Cristian <[email protected]>
AuthorDate: Thu Jul 16 01:57:27 2026 +0300
[Fix][Relax][ONNX] Relax op normalization for onnx subgraphs (#20010)
### Summary
Onnx subgraph imports should also normalize and generate ty_info for its
ops, this is broken since #19853 refactor.
### Issue
```
tests/python/relax/test_frontend_onnx.py:11581: in test_if_subgraph
tvm_model = from_onnx(model, keep_params_in_input=True)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
/usr/lib/python3.15/site-packages/tvm/relax/frontend/onnx/onnx_frontend.py:6283:
in from_onnx
return g.from_onnx(graph, opset)
^^^^^^^^^^^^^^^^^^^^^^^^^
/usr/lib/python3.15/site-packages/tvm/relax/frontend/onnx/onnx_frontend.py:5823:
in from_onnx
self._construct_nodes(graph)
/usr/lib/python3.15/site-packages/tvm/relax/frontend/onnx/onnx_frontend.py:5969:
in _construct_nodes
then_expr = self._convert_subgraph(self.bb, attr["then_branch"])
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
/usr/lib/python3.15/site-packages/tvm/relax/frontend/onnx/onnx_frontend.py:6166:
in _convert_subgraph
op = self._convert_operator(op_name, inputs, attr, self.opset)
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
/usr/lib/python3.15/site-packages/tvm/relax/frontend/onnx/onnx_frontend.py:6117:
in _convert_operator
sym = op_function(self.bb, inputs, attrs, [self._nodes, self._params])
^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
/usr/lib/python3.15/site-packages/tvm/relax/frontend/onnx/onnx_frontend.py:1882:
in _impl_v11
ndim = len(inputs[0].ty.shape)
^^^^^^^^^^^^^^^^^^
E AttributeError: 'Type' object has no attribute 'shape'
```
### Fix
Add conversion check, normalize and populate the final relax op with
ty_info regardless of the graph context.
---
python/tvm/relax/frontend/onnx/onnx_frontend.py | 14 +--
tests/python/relax/test_frontend_onnx.py | 115 ++++++++++++++++++++++++
2 files changed, 118 insertions(+), 11 deletions(-)
diff --git a/python/tvm/relax/frontend/onnx/onnx_frontend.py
b/python/tvm/relax/frontend/onnx/onnx_frontend.py
index 45cd747e98..eaad127524 100644
--- a/python/tvm/relax/frontend/onnx/onnx_frontend.py
+++ b/python/tvm/relax/frontend/onnx/onnx_frontend.py
@@ -6015,9 +6015,6 @@ class ONNXGraphImporter:
raise ValueError(f"Node {node.name} cannot handle
ShapeExpr inputs.")
try:
op = self._convert_operator(op_name, inputs, attr, self.opset)
- # Create type information for the new operator.
- if isinstance(op, relax.Expr):
- op = self.bb.normalize(op)
except Exception as err: # pylint: disable=broad-exception-caught
print(f"Error converting operator {op_name}, with inputs:
{inputs}")
raise err
@@ -6117,6 +6114,9 @@ class ONNXGraphImporter:
sym = op_function(self.bb, inputs, attrs, [self._nodes,
self._params])
else:
raise NotImplementedError(f"Operator {op_name} not implemented.")
+ # Create type information for the new operator.
+ if isinstance(sym, relax.Expr):
+ sym = self.bb.normalize(sym)
return sym
def _convert_subgraph(self, bb, graph):
@@ -6164,14 +6164,6 @@ class ONNXGraphImporter:
continue
op = self._convert_operator(op_name, inputs, attr, self.opset)
- try:
- _ = op.ty
- has_ty = True
- except tvm.error.InternalError:
- has_ty = False
-
- if not has_ty:
- op = bb.normalize(op)
if not isinstance(op, relax.Tuple):
if isinstance(op.ty, relax.TupleType):
diff --git a/tests/python/relax/test_frontend_onnx.py
b/tests/python/relax/test_frontend_onnx.py
index 7269594104..c0f860aaf3 100644
--- a/tests/python/relax/test_frontend_onnx.py
+++ b/tests/python/relax/test_frontend_onnx.py
@@ -11533,6 +11533,121 @@ def test_if_nested():
tvm.ir.assert_structural_equal(tvm_model, Expected)
+def test_if_subgraph():
+ """Test If subgraph."""
+ input_tensor_info = helper.make_tensor_value_info("input",
TensorProto.FLOAT, [1, 3, 3])
+ cond_tensor_info = helper.make_tensor_value_info("cond", TensorProto.BOOL,
[])
+ y_tensor_info = helper.make_tensor_value_info("Y", TensorProto.FLOAT, [1,
1, 3, 3])
+ b_tensor_info = helper.make_tensor_value_info("B", TensorProto.FLOAT, [1,
1, 3, 3])
+ c_tensor_info = helper.make_tensor_value_info("C", TensorProto.FLOAT, [1,
1, 3, 3])
+ unsqueeze_axes_tensor = helper.make_tensor(
+ name="unsqueeze_axes", data_type=TensorProto.INT64, dims=[1], vals=[0]
+ )
+ unsqueeze_then_node = helper.make_node(
+ "Unsqueeze", inputs=["input", "unsqueeze_axes"],
outputs=["input_unsqueezed_then"]
+ )
+ then_out_info = helper.make_tensor_value_info("then_out",
TensorProto.FLOAT, [1, 1, 3, 3])
+ then_node = helper.make_node(
+ "Conv",
+ inputs=["input_unsqueezed_then", "B"],
+ outputs=["then_out"],
+ dilations=[1, 1],
+ group=1,
+ kernel_shape=[3, 3],
+ pads=[1, 1, 1, 1],
+ strides=[1, 1],
+ )
+ then_graph = helper.make_graph(
+ nodes=[unsqueeze_then_node, then_node],
+ name="then_branch_graph",
+ inputs=[],
+ outputs=[then_out_info],
+ )
+ unsqueeze_else_node = helper.make_node(
+ "Unsqueeze", inputs=["input", "unsqueeze_axes"],
outputs=["input_unsqueezed_else"]
+ )
+ else_out_info = helper.make_tensor_value_info("else_out",
TensorProto.FLOAT, [1, 1, 3, 3])
+ else_node = helper.make_node(
+ "Conv",
+ inputs=["input_unsqueezed_else", "C"],
+ outputs=["else_out"],
+ dilations=[1, 1],
+ group=1,
+ kernel_shape=[3, 3],
+ pads=[1, 1, 1, 1],
+ strides=[1, 1],
+ )
+ else_graph = helper.make_graph(
+ nodes=[unsqueeze_else_node, else_node],
+ name="else_branch_graph",
+ inputs=[],
+ outputs=[else_out_info],
+ )
+
+ if_node = helper.make_node(
+ "If", inputs=["cond"], outputs=["Y"], then_branch=then_graph,
else_branch=else_graph
+ )
+ outer_graph = helper.make_graph(
+ nodes=[if_node],
+ name="CondSubgraph",
+ inputs=[cond_tensor_info, input_tensor_info, b_tensor_info,
c_tensor_info],
+ outputs=[y_tensor_info],
+ initializer=[unsqueeze_axes_tensor],
+ )
+ opset_imports = [helper.make_operatorsetid("", 15)]
+ model = helper.make_model(
+ outer_graph, producer_name="condsubgraph", opset_imports=opset_imports
+ )
+
+ tvm_model = from_onnx(model, keep_params_in_input=True)
+ tvm_model, _ = tvm.relax.frontend.detach_params(tvm_model)
+
+ @I.ir_module
+ class Expected:
+ @R.function
+ def main(
+ cond: R.Tensor((), dtype="bool"),
+ input: R.Tensor((1, 3, 3), dtype="float32"),
+ B: R.Tensor((1, 1, 3, 3), dtype="float32"),
+ C: R.Tensor((1, 1, 3, 3), dtype="float32"),
+ unsqueeze_axes: R.Tensor((1,), dtype="int64"),
+ ) -> R.Tensor((1, 1, 3, 3), dtype="float32"):
+ R.func_attr({"num_input": 4})
+ gv: R.Tensor((1, 1, 3, 3), dtype="float32") = R.expand_dims(input,
axis=[0])
+ gv1: R.Tensor((1, 1, 3, 3), dtype="float32") =
R.expand_dims(input, axis=[0])
+ if cond:
+ gv2: R.Tensor((1, 1, 3, 3), dtype="float32") = R.nn.conv2d(
+ gv,
+ B,
+ strides=[1, 1],
+ padding=[1, 1, 1, 1],
+ dilation=[1, 1],
+ groups=1,
+ data_layout="NCHW",
+ kernel_layout="OIHW",
+ out_layout="NCHW",
+ out_dtype=None,
+ )
+ gv4: R.Tensor((1, 1, 3, 3), dtype="float32") = gv2
+ else:
+ gv3: R.Tensor((1, 1, 3, 3), dtype="float32") = R.nn.conv2d(
+ gv1,
+ C,
+ strides=[1, 1],
+ padding=[1, 1, 1, 1],
+ dilation=[1, 1],
+ groups=1,
+ data_layout="NCHW",
+ kernel_layout="OIHW",
+ out_layout="NCHW",
+ out_dtype=None,
+ )
+ gv4: R.Tensor((1, 1, 3, 3), dtype="float32") = gv3
+ return gv4
+
+ tvm.ir.assert_structural_equal(tvm_model, Expected)
+
+
# Helper that builds the ONNX graph for MatMulInteger so the tests don't
repeat boilerplate code every time
def _make_matmulinteger_model(A_shape, B_shape, A_dtype, B_dtype,
a_zp_array=None, b_zp_array=None):
"""Build a minimal single-node ONNX graph for MatMulInteger."""