Skip to content

Commit 0e71891

Browse files
authored
Fix Slice op's shape inference logic (#2526)
* Fix slice shape inference * Fix build
1 parent 1211141 commit 0e71891

2 files changed

Lines changed: 20 additions & 1 deletion

File tree

‎onnx/defs/tensor/defs.cc‎

Lines changed: 9 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -679,8 +679,16 @@ ONNX_OPERATOR_SET_SCHEMA(
679679

680680
// input dim value is missing - cannot perform shape inference for
681681
// this axis
682-
if (!input_dim.has_dim_value())
682+
if (!input_dim.has_dim_value()) {
683+
// Clear any previously propagated dim_param and leave this dimension "empty",
684+
// before moving on to the next dimension
685+
ctx.getOutputType(0)
686+
->mutable_tensor_type()
687+
->mutable_shape()
688+
->mutable_dim(static_cast<int>(axis))
689+
->clear_dim_param();
683690
continue;
691+
}
684692

685693
const auto input_dim_value = input_dim.dim_value();
686694

‎onnx/test/shape_inference_test.py‎

Lines changed: 11 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -500,6 +500,17 @@ def test_slice_with_input_shape(self): # type: () -> None
500500
make_tensor('ends', TensorProto.INT64, (2, ), (2, 2))])
501501
self._assert_inferred(graph, [make_tensor_value_info('y', TensorProto.FLOAT, (1, 2))])
502502

503+
def test_slice_with_input_shape_containing_dim_params(self): # type: () -> None
504+
graph = self._make_graph(
505+
[('x', TensorProto.FLOAT, (1, 'a', 1)),
506+
('starts', TensorProto.INT64, (3,)),
507+
('ends', TensorProto.INT64, (3,))],
508+
[make_node('Slice', ['x', 'starts', 'ends'], ['y'])],
509+
[],
510+
initializer=[make_tensor('starts', TensorProto.INT64, (3,), (0, 0, 0)),
511+
make_tensor('ends', TensorProto.INT64, (3,), (1, 1, 1))])
512+
self._assert_inferred(graph, [make_tensor_value_info('y', TensorProto.FLOAT, (1, None, 1))]) # type: ignore
513+
503514
def test_slice_with_input_shape_steps(self): # type: () -> None
504515
graph = self._make_graph(
505516
[('x', TensorProto.FLOAT, (5, 6, 7)),

0 commit comments

Comments
 (0)