From 8a98a947cfcc7cf3a6df415469db5cce208ea173 Mon Sep 17 00:00:00 2001 From: Javier de Jesus Date: Sat, 3 Oct 2026 13:52:02 +0000 Subject: [PATCH] [Fix][Relax][ONNX] Support runtime shape tensors in Reshape The Relax ONNX importer handled `Reshape` shape inputs that are constants or `Shape` outputs. When the shape is a runtime int64 tensor, such as a graph input, the converter passed the tensor to `relax.op.reshape`, which fails with `TypeError: Reshape requires the input new shape to be Shape. However, the given one is relax.TensorType`. For a runtime shape, the converter now builds the target shape with Relax tensor ops and converts it to a `ShapeExpr` with `_tensor_to_shape_expr`, the way the dynamic-axes path of `Unsqueeze` does. With `allowzero=0` a 0 entry copies the input dimension at the same index, and a -1 entry is the number of input elements divided by the product of the other entries, so neither reaches `relax.op.reshape` as a literal. The copy index is clamped to the last input dimension so `take` stays in bounds, and a scalar input skips the copy. With `allowzero=1` a literal 0 stays 0. The constant and `ShapeExpr` paths are unchanged. The new tests make the shape a graph input and compare the imported model with onnxruntime for the shape from the issue, `-1`, `0`, a rank increase and decrease, a scalar input, a symbolic batch dimension, and `allowzero=1` including a zero-size input. One more test checks the error for an unknown shape length, and one checks that a `Shape` output of an unknown-rank input still imports. The 16 runtime-shape tests fail on `main` with the `TypeError` above. Fixes #20174 Related to #17892, which reports the same error on v0.20.0. Tests: - `pytest -q -n 6 tests/python/relax/test_frontend_onnx.py`: 633 passed, 7 skipped, 4 xfailed - `pre-commit run --files` on both changed files: passed --- .../tvm/relax/frontend/onnx/onnx_frontend.py | 31 +++++++ tests/python/relax/test_frontend_onnx.py | 85 +++++++++++++++++++ 2 files changed, 116 insertions(+) diff --git a/python/tvm/relax/frontend/onnx/onnx_frontend.py b/python/tvm/relax/frontend/onnx/onnx_frontend.py index c078073c3383..042d399c9e89 100644 --- a/python/tvm/relax/frontend/onnx/onnx_frontend.py +++ b/python/tvm/relax/frontend/onnx/onnx_frontend.py @@ -1759,6 +1759,37 @@ def _impl_v13(cls, bb, inputs, attr, params): ) else: new_shape = new_shape_values + elif isinstance(getattr(new_shape, "ty", None), relax.TensorType): + new_shape = _as_int64_tensor(bb, new_shape) + shape_len = _get_known_tensor_length(new_shape) + if shape_len is None: + raise ValueError("Reshape requires a statically known shape length.") + data_ndim = _get_known_tensor_rank(data) + if data_ndim is None: + raise ValueError("Reshape requires a statically known input rank.") + data_dims = bb.normalize(relax.op.shape_to_tensor(relax.op.shape_of(data))) + if not allowzero and data_ndim > 0: + # A 0 copies the input dimension at the same index. + copy_indices = relax.op.minimum( + relax.op.arange(shape_len, dtype="int64"), relax.const(data_ndim - 1, "int64") + ) + copied_dims = relax.op.take(data_dims, copy_indices, axis=0) + new_shape = bb.normalize( + relax.op.where( + relax.op.equal(new_shape, relax.const(0, "int64")), copied_dims, new_shape + ) + ) + # A -1 is the number of elements divided by the product of the other dims. + is_inferred = relax.op.equal(new_shape, relax.const(-1, "int64")) + known_numel = relax.op.prod( + relax.op.where(is_inferred, relax.const(1, "int64"), new_shape), axis=[0] + ) + inferred_dim = relax.op.floor_divide( + relax.op.prod(data_dims, axis=[0]), + relax.op.maximum(known_numel, relax.const(1, "int64")), + ) + new_shape = bb.normalize(relax.op.where(is_inferred, inferred_dim, new_shape)) + new_shape = _tensor_to_shape_expr(bb, new_shape, shape_len, "reshape_dim") 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 df3eb29cef71..e133ef7d617f 100644 --- a/tests/python/relax/test_frontend_onnx.py +++ b/tests/python/relax/test_frontend_onnx.py @@ -3175,6 +3175,91 @@ def main( verify_reshape_shape_output([3, 1], [3, 1], ExpectedRank2ColumnShape) +@pytest.mark.parametrize( + "data_shape, shape, allowzero, symbolic_batch, out_shape", + [ + ([2, 3], [3, 2], 0, False, [3, 2]), + ([2, 3], [-1, 2], 0, False, [3, 2]), + ([2, 3], [0, 3], 0, False, [2, 3]), + ([2, 3, 4], [0, -1], 0, False, [2, 12]), + ([2, 3, 4], [0, 0, -1], 0, False, [2, 3, 4]), + ([6], [1, 2, 3], 0, False, [1, 2, 3]), + ([2, 3, 4], [24], 0, False, [24]), + ([2, 3, 4], [-1], 0, False, [24]), + ([], [1], 0, False, [1]), + ([2, 3], [-1, 3], 0, True, [2, 3]), + ([2, 3], [0, -1], 0, True, [2, 3]), + ([3, 4], [-1, 2], 1, False, [6, 2]), + ([3, 4], [2, 6], 1, False, [2, 6]), + ([2, 0], [0, 2], 1, False, [0, 2]), + ([2, 3, 4], [-1, 12], 1, True, [2, 12]), + ], +) +def test_reshape_runtime_shape(data_shape, shape, allowzero, symbolic_batch, out_shape): + """The shape input is a graph input, so its values are only known at run time.""" + reshape_node = helper.make_node("Reshape", ["data", "shape"], ["reshaped"], allowzero=allowzero) + declared_shape = ["batch", *data_shape[1:]] if symbolic_batch else data_shape + graph = helper.make_graph( + [reshape_node], + "reshape_runtime_shape_test", + inputs=[ + helper.make_tensor_value_info("data", TensorProto.FLOAT, declared_shape), + helper.make_tensor_value_info("shape", TensorProto.INT64, [len(shape)]), + ], + outputs=[helper.make_tensor_value_info("reshaped", TensorProto.FLOAT, out_shape)], + ) + model = helper.make_model(graph, producer_name="reshape_runtime_shape_test") + inputs = { + "data": np.arange(int(np.prod(data_shape)), dtype="float32").reshape(data_shape), + "shape": np.array(shape, dtype="int64"), + } + + check_correctness(model, inputs=inputs, opset=14) + + +def test_reshape_runtime_shape_unknown_length(): + reshape_node = helper.make_node("Reshape", ["data", "shape"], ["reshaped"]) + graph = helper.make_graph( + [reshape_node], + "reshape_runtime_shape_unknown_length_test", + inputs=[ + helper.make_tensor_value_info("data", TensorProto.FLOAT, [2, 3]), + helper.make_tensor_value_info("shape", TensorProto.INT64, ["length"]), + ], + outputs=[helper.make_tensor_value_info("reshaped", TensorProto.FLOAT, [2, 3])], + ) + model = helper.make_model( + graph, + producer_name="reshape_runtime_shape_unknown_length_test", + opset_imports=[helper.make_opsetid("", 14)], + ) + + with pytest.raises(ValueError, match="Reshape requires a statically known shape length"): + from_onnx(model, opset=14, keep_params_in_input=True) + + +def test_reshape_shape_typed_runtime_shape(): + """A Shape-typed value (not a tensor) as the target shape keeps working.""" + nodes = [ + helper.make_node("Shape", ["like"], ["like_shape"]), + helper.make_node("Reshape", ["data", "like_shape"], ["reshaped"]), + ] + graph = helper.make_graph( + nodes, + "reshape_shape_typed_runtime_shape_test", + inputs=[ + helper.make_tensor_value_info("like", TensorProto.FLOAT, None), + helper.make_tensor_value_info("data", TensorProto.FLOAT, [2, 12]), + ], + outputs=[helper.make_tensor_value_info("reshaped", TensorProto.FLOAT, [4, 3, 2])], + ) + model = helper.make_model(graph, producer_name="reshape_shape_typed_runtime_shape_test") + + tvm_model = from_onnx(model, opset=13, keep_params_in_input=True) + + assert "reshape" in tvm_model.script() + + def test_transpose_scalar(): """Test Transpose with scalar inputs - should return scalar unchanged.""" scalar_node = helper.make_node("Transpose", ["x"], ["y"])