From edd3f4bae539093ec969a8deb21516fd863bd3f3 Mon Sep 17 00:00:00 2001 From: tandede <1090179959@qq.com> Date: Thu, 20 Aug 2026 02:45:57 +0800 Subject: [PATCH 1/2] [Fix][Relax] Honor ONNX Reshape zero semantics ONNX Reshape copies the corresponding input dimension for zero entries unless allowzero is enabled. Normalize those entries before constant folding, and materialize allowzero shapes without triggering Relax reshape zero-copy semantics. Add regression coverage for both cases. --- .../tvm/relax/frontend/onnx/onnx_frontend.py | 19 +++++++- tests/python/relax/test_frontend_onnx.py | 44 +++++++++++++++++++ 2 files changed, 61 insertions(+), 2 deletions(-) diff --git a/python/tvm/relax/frontend/onnx/onnx_frontend.py b/python/tvm/relax/frontend/onnx/onnx_frontend.py index 65bd5bfe1a2f..997319c39362 100644 --- a/python/tvm/relax/frontend/onnx/onnx_frontend.py +++ b/python/tvm/relax/frontend/onnx/onnx_frontend.py @@ -1561,6 +1561,7 @@ class Reshape(OnnxOpConverter): def _impl_v13(cls, bb, inputs, attr, params): data = inputs[0] new_shape = get_constant(inputs[1], params) + allowzero = attr.get("allowzero", 0) if isinstance(data, relax.ShapeExpr): # Preserve identity flatten for shape values to keep shape-specialized @@ -1574,9 +1575,23 @@ def _impl_v13(cls, bb, inputs, attr, params): data = bb.normalize(relax.op.shape_to_tensor(data)) if isinstance(data, relax.Constant) and isinstance(new_shape, relax.Constant): - out = _np.reshape(data.data.numpy(), new_shape.data.numpy().tolist()) + data_array = data.data.numpy() + new_shape_values = new_shape.data.numpy().tolist() + if not allowzero: + new_shape_values = [ + data_array.shape[i] if dim == 0 else dim + for i, dim in enumerate(new_shape_values) + ] + out = _np.reshape(data_array, new_shape_values) return relax.const(out, out.dtype) - if isinstance(new_shape, relax.Constant): + if allowzero: + if isinstance(new_shape, relax.ShapeExpr): + new_shape = bb.normalize(relax.op.shape_to_tensor(new_shape)) + new_shape_ndim = _get_known_tensor_length(new_shape) + if new_shape_ndim is None: + raise ValueError("Reshape requires a statically known output rank.") + new_shape = _tensor_to_shape_expr(bb, new_shape, new_shape_ndim, "reshape_dim") + elif isinstance(new_shape, relax.Constant): new_shape = new_shape.data.numpy().tolist() out = relax.op.reshape(data, new_shape) return out diff --git a/tests/python/relax/test_frontend_onnx.py b/tests/python/relax/test_frontend_onnx.py index 673c25bcddd8..74aa6bbe15b6 100644 --- a/tests/python/relax/test_frontend_onnx.py +++ b/tests/python/relax/test_frontend_onnx.py @@ -2274,6 +2274,50 @@ def main( verify_reshape([7, 32, 32, 8], [0, 32, 32, 8], [7, 32, 32, 8], ExpectedCopyInputDim) +def test_reshape_constant_zero_copy_dimension(): + data = np.arange(6, dtype="float32").reshape(2, 3) + reshape_node = helper.make_node("Reshape", ["data", "shape"], ["reshaped"]) + graph = helper.make_graph( + [reshape_node], + "reshape_constant_zero_copy_test", + inputs=[], + initializer=[ + numpy_helper.from_array(data, "data"), + helper.make_tensor("shape", TensorProto.INT64, [2], [0, 3]), + ], + outputs=[helper.make_tensor_value_info("reshaped", TensorProto.FLOAT, [2, 3])], + ) + model = helper.make_model( + graph, + producer_name="reshape_constant_zero_copy_test", + opset_imports=[helper.make_opsetid("", 14)], + ) + + output = run_in_tvm(model, opset=14) + + tvm.testing.assert_allclose(output.numpy(), data) + + +def test_reshape_allowzero_literal_zero_dimension(): + reshape_node = helper.make_node("Reshape", ["data", "shape"], ["reshaped"], allowzero=1) + graph = helper.make_graph( + [reshape_node], + "reshape_allowzero_test", + inputs=[helper.make_tensor_value_info("data", TensorProto.FLOAT, [2, 0])], + initializer=[helper.make_tensor("shape", TensorProto.INT64, [2], [0, 2])], + outputs=[helper.make_tensor_value_info("reshaped", TensorProto.FLOAT, [0, 2])], + ) + model = helper.make_model( + graph, + producer_name="reshape_allowzero_test", + opset_imports=[helper.make_opsetid("", 14)], + ) + + output = run_in_tvm(model, {"data": np.zeros((2, 0), dtype="float32")}, opset=14) + + assert tuple(output.shape) == (0, 2) + + def test_reshape_shape_output(): def verify_reshape_shape_output(target_shape, output_shape, expected): shape_node = helper.make_node("Shape", ["data"], ["shape_out"]) From d20b94eb4c66c7f5df45d410f8e48b8bfc487add Mon Sep 17 00:00:00 2001 From: tandede <1090179959@qq.com> Date: Sun, 23 Aug 2026 12:54:06 +0800 Subject: [PATCH 2/2] [Fix][Relax] Preserve reshape inference with allowzero --- .../tvm/relax/frontend/onnx/onnx_frontend.py | 17 +++++++-------- tests/python/relax/test_frontend_onnx.py | 21 +++++++++++++++++++ 2 files changed, 29 insertions(+), 9 deletions(-) diff --git a/python/tvm/relax/frontend/onnx/onnx_frontend.py b/python/tvm/relax/frontend/onnx/onnx_frontend.py index 997319c39362..7e8616f65f7b 100644 --- a/python/tvm/relax/frontend/onnx/onnx_frontend.py +++ b/python/tvm/relax/frontend/onnx/onnx_frontend.py @@ -1584,15 +1584,14 @@ def _impl_v13(cls, bb, inputs, attr, params): ] out = _np.reshape(data_array, new_shape_values) return relax.const(out, out.dtype) - if allowzero: - if isinstance(new_shape, relax.ShapeExpr): - new_shape = bb.normalize(relax.op.shape_to_tensor(new_shape)) - new_shape_ndim = _get_known_tensor_length(new_shape) - if new_shape_ndim is None: - raise ValueError("Reshape requires a statically known output rank.") - new_shape = _tensor_to_shape_expr(bb, new_shape, new_shape_ndim, "reshape_dim") - elif isinstance(new_shape, relax.Constant): - new_shape = new_shape.data.numpy().tolist() + if isinstance(new_shape, relax.Constant): + new_shape_values = new_shape.data.numpy().tolist() + if allowzero and 0 in new_shape_values: + new_shape = _tensor_to_shape_expr( + bb, new_shape, len(new_shape_values), "reshape_dim" + ) + else: + new_shape = new_shape_values out = relax.op.reshape(data, new_shape) return out diff --git a/tests/python/relax/test_frontend_onnx.py b/tests/python/relax/test_frontend_onnx.py index 74aa6bbe15b6..09f3a7b01d88 100644 --- a/tests/python/relax/test_frontend_onnx.py +++ b/tests/python/relax/test_frontend_onnx.py @@ -2318,6 +2318,27 @@ def test_reshape_allowzero_literal_zero_dimension(): assert tuple(output.shape) == (0, 2) +def test_reshape_allowzero_infers_dimension_without_literal_zero(): + reshape_node = helper.make_node("Reshape", ["data", "shape"], ["reshaped"], allowzero=1) + graph = helper.make_graph( + [reshape_node], + "reshape_allowzero_infer_dimension_test", + inputs=[helper.make_tensor_value_info("data", TensorProto.FLOAT, [3, 4])], + initializer=[helper.make_tensor("shape", TensorProto.INT64, [2], [-1, 2])], + outputs=[helper.make_tensor_value_info("reshaped", TensorProto.FLOAT, [6, 2])], + ) + model = helper.make_model( + graph, + producer_name="reshape_allowzero_infer_dimension_test", + opset_imports=[helper.make_opsetid("", 14)], + ) + data = np.arange(12, dtype="float32").reshape(3, 4) + + output = run_in_tvm(model, {"data": data}, opset=14) + + tvm.testing.assert_allclose(output.numpy(), data.reshape(6, 2)) + + def test_reshape_shape_output(): def verify_reshape_shape_output(target_shape, output_shape, expected): shape_node = helper.make_node("Shape", ["data"], ["shape_out"])