1 parent 1211141 commit 0e71891Copy full SHA for 0e71891
2 files changed
onnx/defs/tensor/defs.cc
@@ -679,8 +679,16 @@ ONNX_OPERATOR_SET_SCHEMA(
679
680
// input dim value is missing - cannot perform shape inference for
681
// this axis
682
- if (!input_dim.has_dim_value())
+ 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();
690
continue;
691
+ }
692
693
const auto input_dim_value = input_dim.dim_value();
694
onnx/test/shape_inference_test.py
@@ -500,6 +500,17 @@ def test_slice_with_input_shape(self): # type: () -> None
500
make_tensor('ends', TensorProto.INT64, (2, ), (2, 2))])
501
self._assert_inferred(graph, [make_tensor_value_info('y', TensorProto.FLOAT, (1, 2))])
502
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
+
514
def test_slice_with_input_shape_steps(self): # type: () -> None
515
graph = self._make_graph(
516
[('x', TensorProto.FLOAT, (5, 6, 7)),
0 commit comments