From ed0261ce0d2ce4a64a90e2178217868cb813d052 Mon Sep 17 00:00:00 2001 From: OmarAzizi Date: Wed, 26 Aug 2026 02:49:05 +0300 Subject: [PATCH 1/2] [Fix][Relax] Raise error on non-unit dim ONNX Squeeze axis --- .../tvm/relax/frontend/onnx/onnx_frontend.py | 18 ++++++++++++++ tests/python/relax/test_frontend_onnx.py | 24 +++++++++++++++++++ 2 files changed, 42 insertions(+) diff --git a/python/tvm/relax/frontend/onnx/onnx_frontend.py b/python/tvm/relax/frontend/onnx/onnx_frontend.py index a6a2a3152e80..a7982647c970 100644 --- a/python/tvm/relax/frontend/onnx/onnx_frontend.py +++ b/python/tvm/relax/frontend/onnx/onnx_frontend.py @@ -2173,6 +2173,23 @@ def _impl_v13(cls, bb, inputs, attr, params): axis = tuple([int(x) for x in axis.data.numpy()]) return cls._squeeze(bb, data, axis) + @classmethod + def _check_squeeze_axes_are_unit_dims(cls, data, axes): + """Raise if any axis to be squeezed has a statically known extent other than 1.""" + rank = _get_known_tensor_rank(data) + if rank is None: + return + ty = data.ty + if not (isinstance(ty, relax.TensorType) and isinstance(ty.shape, relax.ShapeExpr)): + return + for axis in _normalize_constant_axes(list(axes), rank, "Squeeze"): + extent = ty.shape.values[axis] + if isinstance(extent, tirx.IntImm) and int(extent.value) != 1: + raise ValueError( + f"Squeeze axis {axis} has size {int(extent.value)}, but only " + "axes of size 1 can be squeezed." + ) + @classmethod def _squeeze(cls, bb, data, axis): # If data is constant, perform computation directly. @@ -2201,6 +2218,7 @@ def _squeeze(cls, bb, data, axis): return relax.op.squeeze(data) if isinstance(axis, tuple): + cls._check_squeeze_axes_are_unit_dims(data, axis) return relax.op.squeeze(data, list(axis)) data_ndim = _get_known_tensor_rank(data) diff --git a/tests/python/relax/test_frontend_onnx.py b/tests/python/relax/test_frontend_onnx.py index 4654a4082a2e..d98490f368e9 100644 --- a/tests/python/relax/test_frontend_onnx.py +++ b/tests/python/relax/test_frontend_onnx.py @@ -4434,6 +4434,30 @@ def main(x: R.Tensor((1, 32, 1, 32), dtype="float32")) -> R.Tensor( tvm.ir.assert_structural_equal(tvm_model, Expected) +def test_squeeze_non_unit_axis_raises(): + # Per the ONNX spec, squeezing an axis whose length is not 1 must raise an error + squeeze_node = helper.make_node("Squeeze", ["x", "axes"], ["y"]) + shape = [2, 3] + + graph = helper.make_graph( + [squeeze_node], + "squeeze_non_unit_axis_test", + inputs=[ + helper.make_tensor_value_info("x", TensorProto.FLOAT, shape), + ], + initializer=[helper.make_tensor("axes", TensorProto.INT64, [1], [0])], + outputs=[helper.make_tensor_value_info("y", TensorProto.FLOAT, [3])], + ) + + model = helper.make_model( + graph, + producer_name="squeeze_non_unit_axis_test", + opset_imports=[helper.make_opsetid("", 13)], + ) + with pytest.raises(ValueError, match="has size 2"): + from_onnx(model, opset=13, keep_params_in_input=True) + + def test_squeeze_constant(): def verify_squeeze_constant(axis, expected): shape = [1, 2, 1, 3] From e8afc7e0526612ef2ca8fadadd9ddf062b28ad76 Mon Sep 17 00:00:00 2001 From: OmarAzizi Date: Thu, 27 Aug 2026 13:23:53 +0300 Subject: [PATCH 2/2] [Fix][Relax] Reject import on symbolic dimensions --- .../tvm/relax/frontend/onnx/onnx_frontend.py | 7 +++++- tests/python/relax/test_frontend_onnx.py | 24 +++++++++++++++++++ 2 files changed, 30 insertions(+), 1 deletion(-) diff --git a/python/tvm/relax/frontend/onnx/onnx_frontend.py b/python/tvm/relax/frontend/onnx/onnx_frontend.py index a7982647c970..feb92fdd8159 100644 --- a/python/tvm/relax/frontend/onnx/onnx_frontend.py +++ b/python/tvm/relax/frontend/onnx/onnx_frontend.py @@ -2184,7 +2184,12 @@ def _check_squeeze_axes_are_unit_dims(cls, data, axes): return for axis in _normalize_constant_axes(list(axes), rank, "Squeeze"): extent = ty.shape.values[axis] - if isinstance(extent, tirx.IntImm) and int(extent.value) != 1: + if not isinstance(extent, tirx.IntImm): + raise ValueError( + f"Squeeze axis {axis} has a symbolic extent that cannot be proven to be " + "1 at import time; only statically known unit-size axes can be squeezed." + ) + if int(extent.value) != 1: raise ValueError( f"Squeeze axis {axis} has size {int(extent.value)}, but only " "axes of size 1 can be squeezed." diff --git a/tests/python/relax/test_frontend_onnx.py b/tests/python/relax/test_frontend_onnx.py index d98490f368e9..cd0dda38ce95 100644 --- a/tests/python/relax/test_frontend_onnx.py +++ b/tests/python/relax/test_frontend_onnx.py @@ -4458,6 +4458,30 @@ def test_squeeze_non_unit_axis_raises(): from_onnx(model, opset=13, keep_params_in_input=True) +def test_squeeze_symbolic_axis_raises(): + # Squeezing an axis whose extent is symbolic can't be proven to be 1 at import time + squeeze_node = helper.make_node("Squeeze", ["x", "axes"], ["y"]) + shape = ["N", 3] + + graph = helper.make_graph( + [squeeze_node], + "squeeze_symbolic_axis_test", + inputs=[ + helper.make_tensor_value_info("x", TensorProto.FLOAT, shape), + ], + initializer=[helper.make_tensor("axes", TensorProto.INT64, [1], [0])], + outputs=[helper.make_tensor_value_info("y", TensorProto.FLOAT, [3])], + ) + + model = helper.make_model( + graph, + producer_name="squeeze_symbolic_axis_test", + opset_imports=[helper.make_opsetid("", 13)], + ) + with pytest.raises(ValueError, match="symbolic extent"): + from_onnx(model, opset=13, keep_params_in_input=True) + + def test_squeeze_constant(): def verify_squeeze_constant(axis, expected): shape = [1, 2, 1, 3]