diff --git a/ngraph/core/src/node.cpp b/ngraph/core/src/node.cpp index 60cc0f7ccc8..c04fb69c65d 100644 --- a/ngraph/core/src/node.cpp +++ b/ngraph/core/src/node.cpp @@ -155,12 +155,13 @@ void ov::Node::set_arguments(const NodeVector& arguments) { } void ov::Node::set_arguments(const OutputVector& arguments) { + // Remove existing inputs of this node + m_inputs.clear(); + // Add this node as a user of each argument. size_t i = 0; for (auto& output : arguments) { - auto output_node = output.get_node(); - auto& output_descriptor = output_node->m_outputs.at(output.get_index()); - m_inputs.emplace_back(this, i++, output_descriptor); + set_argument(i++, output); } } diff --git a/ngraph/test/node_input_output.cpp b/ngraph/test/node_input_output.cpp index 0b99f8fb207..ec4c1182e48 100644 --- a/ngraph/test/node_input_output.cpp +++ b/ngraph/test/node_input_output.cpp @@ -96,3 +96,27 @@ TEST(node_input_output, output_create_const) { EXPECT_THROW(add->output(1), std::out_of_range); } + +TEST(node_input_output, input_set_argument) { + auto x = make_shared(element::f32, Shape{1}); + auto y = make_shared(element::f32, Shape{2}); + auto z = make_shared(element::f32, Shape{3}); + + auto add = make_shared(x, y); + + EXPECT_EQ(add->get_input_size(), 2); + EXPECT_EQ(add->input(0).get_shape(), Shape{1}); + EXPECT_EQ(add->input(1).get_shape(), Shape{2}); + + add->set_argument(1, z); + + EXPECT_EQ(add->get_input_size(), 2); + EXPECT_EQ(add->input(0).get_shape(), Shape{1}); + EXPECT_EQ(add->input(1).get_shape(), Shape{3}); + + add->set_arguments(NodeVector{z, x}); + + EXPECT_EQ(add->get_input_size(), 2); + EXPECT_EQ(add->input(0).get_shape(), Shape{3}); + EXPECT_EQ(add->input(1).get_shape(), Shape{1}); +} diff --git a/runtime/bindings/python/src/compatibility/pyngraph/node.cpp b/runtime/bindings/python/src/compatibility/pyngraph/node.cpp index 3c7436d2dc8..ee0a96ee82c 100644 --- a/runtime/bindings/python/src/compatibility/pyngraph/node.cpp +++ b/runtime/bindings/python/src/compatibility/pyngraph/node.cpp @@ -257,6 +257,14 @@ void regclass_pyngraph_Node(py::module m) { Operation version. )"); + node.def("set_argument", &ngraph::Node::set_argument); + node.def("set_arguments", [](const std::shared_ptr& self, const ngraph::NodeVector& args) { + self->set_arguments(args); + }); + node.def("set_arguments", [](const std::shared_ptr& self, const ngraph::OutputVector& args) { + self->set_arguments(args); + }); + node.def_property_readonly("shape", &ngraph::Node::get_shape); node.def_property_readonly("name", &ngraph::Node::get_name); node.def_property_readonly("rt_info", diff --git a/runtime/bindings/python/tests/test_ngraph/test_basic.py b/runtime/bindings/python/tests/test_ngraph/test_basic.py index af6d9bb57ce..2ec656a800b 100644 --- a/runtime/bindings/python/tests/test_ngraph/test_basic.py +++ b/runtime/bindings/python/tests/test_ngraph/test_basic.py @@ -253,6 +253,44 @@ def test_constant_get_data_unsigned_integer(data_type): assert np.allclose(input_data, retrieved_data) +def test_set_argument(): + runtime = get_runtime() + + data1 = np.array([1, 2, 3]) + data2 = np.array([4, 5, 6]) + data3 = np.array([7, 8, 9]) + + node1 = ng.constant(data1, dtype=np.float32) + node2 = ng.constant(data2, dtype=np.float32) + node3 = ng.constant(data3, dtype=np.float32) + node_add = ng.add(node1, node2) + + # Original arguments + computation = runtime.computation(node_add) + output = computation() + assert np.allclose(data1 + data2, output) + + # Arguments changed by set_argument + node_add.set_argument(1, node3.output(0)) + output = computation() + assert np.allclose(data1 + data3, output) + + # Arguments changed by set_argument + node_add.set_argument(0, node3.output(0)) + output = computation() + assert np.allclose(data3 + data3, output) + + # Arguments changed by set_argument(OutputVector) + node_add.set_arguments([node2.output(0), node3.output(0)]) + output = computation() + assert np.allclose(data2 + data3, output) + + # Arguments changed by set_arguments(NodeVector) + node_add.set_arguments([node1, node2]) + output = computation() + assert np.allclose(data1 + data2, output) + + def test_result(): node = np.array([[11, 10], [1, 8], [3, 4]]) result = run_op_node([node], ng.result)