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:
parent
277a23b8e4
commit
702633073e
|
|
@ -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);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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});
|
||||
}
|
||||
|
|
|
|||
|
|
@ -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",
|
||||
|
|
|
|||
|
|
@ -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)
|
||||
|
|
|
|||
Loading…
Reference in New Issue