Skip to content
Merged
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: 15 additions & 3 deletions python/tvm/relay/frontend/onnx.py
Original file line number Diff line number Diff line change
Expand Up @@ -826,6 +826,15 @@ def _impl_v1(cls, inputs, attr, params):
return out


def is_ort_version_greater_than(ver):
import onnxruntime as ort

v11, v12, v13 = tuple(int(v) for v in ort.__version__.split("."))
Comment thread
padreofthegame marked this conversation as resolved.
v21, v22, v23 = tuple(int(v) for v in ver.split("."))

return (v11 > v21) or (v11 == v21 and v12 > v22) or ((v11, v12) == (v21, v22) and v13 > v23)


class ConvTranspose(OnnxOpConverter):
"""Operator converter for ConvTranspose."""

Expand Down Expand Up @@ -963,12 +972,15 @@ def _impl_v11(cls, inputs, attr, params):
)
left = [p // 2 for p in total_pad]
right = [total_pad[i] - left[i] for i in range(kndim)]

if "output_shape" in attr and "auto_pad" not in attr:
pad = right + left
elif "LOWER" in attr["auto_pad"]:
pad = left + right
else:
elif ("LOWER" in attr["auto_pad"] and is_ort_version_greater_than("1.12.1")) or (
("UPPER" in attr["auto_pad"] and not is_ort_version_greater_than("1.12.1"))
):
pad = right + left
else:
pad = left + right
attr["pads"] = pad
elif attr["auto_pad"] == "VALID":
attr["pads"] = tuple([0 for i in range(ndim - 2)])
Expand Down
46 changes: 45 additions & 1 deletion tests/python/frontend/onnx/test_forward.py
Original file line number Diff line number Diff line change
Expand Up @@ -3404,6 +3404,36 @@ def repeat(num, dims):
auto_pad="SAME_LOWER",
)

verify_convtranspose_with_output_shape(
(1, 1) + repeat(32, dims),
(1, 2) + repeat(4, dims),
repeat(num, dims),
repeat(4, dims),
repeat(2, dims),
repeat(1, dims),
auto_pad="SAME_UPPER",
)

verify_convtranspose_with_output_shape(
(1, 1, 3, 3),
(1, 2, 3, 3),
(6, 6),
(3, 3),
(2, 2),
(1, 1),
auto_pad="SAME_UPPER",
)

verify_convtranspose_with_output_shape(
(1, 1, 3, 3),
(1, 2, 3, 3),
(6, 6),
(3, 3),
(2, 2),
(1, 1),
auto_pad="SAME_LOWER",
)


@tvm.testing.parametrize_targets
def test_unsqueeze_constant(target, dev):
Expand Down Expand Up @@ -5634,7 +5664,6 @@ def verify_eyelike(indata, dynamic=False):
"test_cast_DOUBLE_to_FLOAT16",
"test_castlike_DOUBLE_to_FLOAT16",
"test_castlike_DOUBLE_to_FLOAT16_expanded",
"test_convtranspose_autopad_same",
"test_convtranspose_dilations",
"test_cumsum_1d",
"test_cumsum_1d_exclusive",
Expand Down Expand Up @@ -5766,6 +5795,15 @@ def _load_proto(proto_filename, target_list, model_type_proto):
)


def is_ort_version_lower_than(ver):
import onnxruntime as ort

v11, v12, v13 = tuple(int(v) for v in ort.__version__.split("."))
v21, v22, v23 = tuple(int(v) for v in ver.split("."))

return (v11 < v21) or (v11 == v21 and v12 < v22) or ((v11, v12) == (v21, v22) and v13 < v23)


@pytest.mark.parametrize("onnx_test", onnx_test_folders)
@tvm.testing.parametrize_targets
def test_onnx_nodes(target, dev, onnx_test):
Expand All @@ -5782,6 +5820,12 @@ def test_onnx_nodes(target, dev, onnx_test):
if onnx_test in target_specific_skips:
pytest.skip(f"Onnx test '{onnx_test}' not yet supported by TVM on {target_kind} targets")

if is_ort_version_lower_than("1.13.1") and onnx_test == "test_convtranspose_autopad_same":
pytest.skip(
f"Onnx test '{onnx_test}' expected to fail for onnxruntime version lower than 1.13.1 "
"due to different interpretation of auto_pad parameters SAME_UPPER and SAME_LOWER."
)

test_dir = os.path.join(onnx_test_node_dir, onnx_test)

atol = 1e-5
Expand Down