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
19 changes: 18 additions & 1 deletion python/tvm/relax/frontend/onnx/onnx_frontend.py
Original file line number Diff line number Diff line change
Expand Up @@ -3719,6 +3719,13 @@ def _impl_v12(cls, bb, inputs, attr, params):
return relax.op.nn.dropout(inputs[0], ratio)


def _none_if_empty_constant(value):
"""Return None for a zero-element constant, which stands for an omitted input."""
if isinstance(value, tvm.ir.GenericConst) and value.value.numpy().size == 0:
return None
return value


def _onnx_resize_spatial_roi_vector(roi_full: relax.Expr, rank: int) -> relax.Expr:
"""Map ONNX ROI [starts..., ends...] to TOPI spatial ROI (drop N/C axes)."""
return relax.op.concat(
Expand Down Expand Up @@ -3846,6 +3853,10 @@ def _impl_v18(cls, bb, inputs, attr, params):
ndims = len(x.ty.shape)
assert ndims in (3, 4, 5), "Only resize1d/resize2d/resize3d are supported."

# Some exporters (e.g. PyTorch at opset <= 12) pass an empty tensor instead of
# omitting an optional input, which the ONNX spec treats as "not provided".
scales = _none_if_empty_constant(scales)
sizes = _none_if_empty_constant(sizes)
assert scales is None or sizes is None, (
"Only one of scales and sizes can be provided in Resize."
)
Expand Down Expand Up @@ -3889,7 +3900,13 @@ def _impl_v18(cls, bb, inputs, attr, params):
if isinstance(sizes, tvm.ir.GenericConst):
sizes = sizes.value.numpy().astype("int64").tolist()[2:]
elif isinstance(sizes, relax.expr.ShapeExpr):
sizes = [int(val.value) for val in sizes.values][2:]
# Sizes computed from Shape/Slice/Concat may carry symbolic batch or
# channel dims; only the spatial dims are needed, and the relax resize
# ops accept symbolic spatial extents.
sizes = [
int(val.value) if isinstance(val, tirx.IntImm) else val
for val in list(sizes.values)[2:]
]
else:
raise ValueError(f"Type {type(sizes)} for size is currently unsupported.")

Expand Down
103 changes: 103 additions & 0 deletions tests/python/relax/test_frontend_onnx.py
Original file line number Diff line number Diff line change
Expand Up @@ -9432,6 +9432,109 @@ def _visit(expr):
assert seen_resize3d


@pytest.mark.parametrize("input_shape", [[1, 3, 4, 5], ["N", 3, 4, 5]])
def test_resize_sizes_with_empty_scales_tensor(input_shape):
# PyTorch (opset <= 12) exports size-based interpolation with `roi` and `scales`
# given as empty constant tensors rather than omitted inputs.
nodes = [
helper.make_node(
"Constant", [], ["roi"], value=helper.make_tensor("", TensorProto.FLOAT, [0], [])
),
helper.make_node(
"Constant", [], ["scales"], value=helper.make_tensor("", TensorProto.FLOAT, [0], [])
),
helper.make_node(
"Resize",
["X", "roi", "scales", "sizes"],
["Y"],
mode="nearest",
coordinate_transformation_mode="asymmetric",
nearest_mode="floor",
),
]
graph = helper.make_graph(
nodes,
"resize_empty_scales",
inputs=[helper.make_tensor_value_info("X", TensorProto.FLOAT, input_shape)],
initializer=[helper.make_tensor("sizes", TensorProto.INT64, [4], [1, 3, 8, 10])],
outputs=[helper.make_tensor_value_info("Y", TensorProto.FLOAT, None)],
)
model = helper.make_model(graph, producer_name="resize_empty_scales")
check_correctness(
model, inputs={"X": generate_random_value([1, 3, 4, 5], TensorProto.FLOAT)}, opset=13
)


def _make_resize_sizes_from_shape_model(input_shape, shape_source_shape=None):
"""Resize whose `sizes` is Concat(Shape(X)[:2], spatial) as exporters emit it.

When `shape_source_shape` is given, the spatial sizes are taken from a second
input's shape instead of constants, so they are symbolic when it is.
"""
inputs = [helper.make_tensor_value_info("X", TensorProto.FLOAT, input_shape)]
nodes = [
helper.make_node("Shape", ["X"], ["x_shape"]),
helper.make_node("Slice", ["x_shape", "zero", "two", "zero"], ["nc"]),
]
initializers = [
helper.make_tensor("zero", TensorProto.INT64, [1], [0]),
helper.make_tensor("two", TensorProto.INT64, [1], [2]),
]
if shape_source_shape is None:
initializers.append(helper.make_tensor("hw", TensorProto.INT64, [2], [8, 10]))
else:
inputs.append(helper.make_tensor_value_info("S", TensorProto.FLOAT, shape_source_shape))
nodes += [
helper.make_node("Shape", ["S"], ["s_shape"]),
helper.make_node("Slice", ["s_shape", "two", "four", "zero"], ["hw"]),
]
initializers.append(helper.make_tensor("four", TensorProto.INT64, [1], [4]))
nodes += [
helper.make_node("Concat", ["nc", "hw"], ["sizes"], axis=0),
helper.make_node(
"Resize",
["X", "", "", "sizes"],
["Y"],
mode="nearest",
coordinate_transformation_mode="asymmetric",
nearest_mode="floor",
),
]
graph = helper.make_graph(
nodes,
"resize_sizes_from_shape",
inputs=inputs,
initializer=initializers,
outputs=[helper.make_tensor_value_info("Y", TensorProto.FLOAT, None)],
)
return helper.make_model(graph, producer_name="resize_sizes_from_shape")


def test_resize_sizes_from_shape_symbolic_batch():
model = _make_resize_sizes_from_shape_model(["N", 3, 4, 5])
func = from_onnx(model, opset=18, keep_params_in_input=True)["main"]
n = func.params[0].ty.shape.values[0]
out_shape = func.ret_ty.shape.values
tvm.ir.assert_structural_equal(out_shape[0], n)
assert [int(v) for v in out_shape[1:]] == [3, 8, 10]

x = generate_random_value([2, 3, 4, 5], TensorProto.FLOAT)
check_correctness(model, inputs={"X": x}, opset=18)


def test_resize_sizes_from_shape_symbolic_spatial():
model = _make_resize_sizes_from_shape_model(["N", 3, 4, 5], ["N", 3, "H", "W"])
func = from_onnx(model, opset=18, keep_params_in_input=True)["main"]
_, _, h, w = func.params[1].ty.shape.values
out_shape = func.ret_ty.shape.values
tvm.ir.assert_structural_equal(out_shape[2], h)
tvm.ir.assert_structural_equal(out_shape[3], w)

x = generate_random_value([2, 3, 4, 5], TensorProto.FLOAT)
s = generate_random_value([2, 3, 8, 10], TensorProto.FLOAT)
check_correctness(model, inputs={"X": x, "S": s}, opset=18)


def test_einsum():
eqn = "ij->i"
einsum_node = helper.make_node("Einsum", ["x"], ["y"], equation=eqn)
Expand Down
Loading