Add set_argument/s methods to nGraph Python API (#7196)

* Add set_argument/s methods to nGraph Python API

* Apply formatting

* Clear inputs in Node::set_arguments

* Run all tests in one container

* Fix formatting

* Add unit test
This commit is contained in:
Michał Karzyński 2021-09-18 05:37:26 +02:00 committed by GitHub
parent 277a23b8e4
commit 702633073e
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
4 changed files with 74 additions and 3 deletions

View File

@ -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);
}
}

View File

@ -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<op::Parameter>(element::f32, Shape{1});
auto y = make_shared<op::Parameter>(element::f32, Shape{2});
auto z = make_shared<op::Parameter>(element::f32, Shape{3});
auto add = make_shared<op::v1::Add>(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});
}

View File

@ -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<ngraph::Node>& self, const ngraph::NodeVector& args) {
self->set_arguments(args);
});
node.def("set_arguments", [](const std::shared_ptr<ngraph::Node>& 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",

View File

@ -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)