Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 16 additions & 2 deletions python/tvm/relax/frontend/onnx/onnx_frontend.py
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand All @@ -1574,10 +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):
new_shape = new_shape.data.numpy().tolist()
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

Expand Down
65 changes: 65 additions & 0 deletions tests/python/relax/test_frontend_onnx.py
Original file line number Diff line number Diff line change
Expand Up @@ -2274,6 +2274,71 @@ 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_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"])
Expand Down