Skip to content

Commit 61f0bbc

Browse files
BowenBaoebarsoum
andauthored
Fix a bug in ScatterND shape inference (#2577)
* Check if input has shape before propagating * add a test case Co-authored-by: Emad Barsoum <[email protected]>
1 parent 05bce9c commit 61f0bbc

2 files changed

Lines changed: 18 additions & 1 deletion

File tree

‎onnx/defs/tensor/defs.cc‎

Lines changed: 3 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -966,7 +966,9 @@ ONNX_OPERATOR_SET_SCHEMA(
966966
"Constrain input and output types to any tensor type.")
967967
.TypeAndShapeInferenceFunction([](InferenceContext& ctx) {
968968
propagateElemTypeFromInputToOutput(ctx, 0, 0);
969-
propagateShapeFromInputToOutput(ctx, 0, 0);
969+
if (hasNInputShapes(ctx, 1)) {
970+
propagateShapeFromInputToOutput(ctx, 0, 0);
971+
}
970972
}));
971973

972974
static const char* ScatterElements_ver11_doc = R"DOC(

‎onnx/test/shape_inference_test.py‎

Lines changed: 15 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -455,6 +455,21 @@ def test_scatternd(self): # type: () -> None
455455
[])
456456
self._assert_inferred(graph, [make_tensor_value_info('y', TensorProto.FLOAT, (4, 5, 6))]) # type: ignore
457457

458+
def test_scatternd_noshape(self): # type: () -> None
459+
# The shape of 'x_reshaped' cannot be inferred, since it is the output of a dynamic reshape.
460+
# Thus the shape of 'y' is also None.
461+
graph = self._make_graph(
462+
[('x', TensorProto.FLOAT, (4, 5, 6)),
463+
('indices', TensorProto.INT64, (3, 3, 2)),
464+
('updates', TensorProto.FLOAT, (3, 3, 6)),
465+
('shape', TensorProto.UNDEFINED, (2,))],
466+
[make_node("Reshape", ['x', 'shape'], ['x_reshaped']),
467+
make_node("ScatterND", ['x_reshaped', 'indices', 'updates'], ['y'])],
468+
[])
469+
self._assert_inferred(graph, [
470+
make_tensor_value_info('x_reshaped', TensorProto.FLOAT, None),
471+
make_tensor_value_info('y', TensorProto.FLOAT, None)]) # type: ignore
472+
458473
def test_squeeze(self): # type: () -> None
459474
graph = self._make_graph(
460475
[('x', TensorProto.FLOAT, (1, 3, 1, 1, 2, 1))],

0 commit comments

Comments
 (0)