diff --git a/tools/mo/openvino/tools/mo/ops/group_norm.py b/tools/mo/openvino/tools/mo/ops/group_norm.py index fa5bb290aaa..24395d5f3d3 100644 --- a/tools/mo/openvino/tools/mo/ops/group_norm.py +++ b/tools/mo/openvino/tools/mo/ops/group_norm.py @@ -18,3 +18,25 @@ class GroupNorm(Op): 'in_ports_count': 3, 'out_ports_count': 1, }, attrs) + + +class GroupNormalization(Op): + op = 'GroupNormalization' + enabled = True + + def __init__(self, graph: Graph, attrs: dict): + super().__init__(graph, { + 'op': self.op, + 'type': self.op, + 'infer': copy_shape_infer, + 'version': 'opset12', + + 'num_groups': None, + 'epsilon': None, + + 'in_ports_count': 3, + 'out_ports_count': 1, + }, attrs) + + def backend_attrs(self): + return ['num_groups', 'epsilon'] diff --git a/tools/mo/openvino/tools/mo/ops/scatter.py b/tools/mo/openvino/tools/mo/ops/scatter.py index c53693d0741..e761d9a35e1 100644 --- a/tools/mo/openvino/tools/mo/ops/scatter.py +++ b/tools/mo/openvino/tools/mo/ops/scatter.py @@ -27,6 +27,9 @@ class Scatter(Op): 'infer': self.infer, 'reverse_infer': lambda node: reverse_bypass_infer(node, in_ports=[0]), + 'reduction': None, + 'use_init_val': None, + 'in_ports_count': 4, 'out_ports_count': 1, } @@ -85,6 +88,13 @@ class ScatterElementsUpdate(Scatter): op = op_type = 'ScatterElementsUpdate' version = 'opset3' + def backend_attrs(self): + version = self.get_opset() + if version == 'opset12': + return ['reduction', 'use_init_val'] + else: + return [] + @staticmethod def infer(node: Node): Scatter.infer(node) @@ -111,7 +121,9 @@ class ScatterElementsUpdate(Scatter): ''.format(node_name, indices_shape, updates_shape) axis = node.in_port(3).data.get_value() - if input_value is not None and indices_value is not None and updates_value is not None and axis is not None: + opset = node.soft_get('version', 'default') + is_opset12_reduction = opset == 'opset12' and (node.soft_get('reduction') != 'none' or not node.soft_get('use_init_val')) + if input_value is not None and indices_value is not None and updates_value is not None and axis is not None and not is_opset12_reduction: assert axis.size == 1, "The node {} has axis input value size equal to {} but it should be exactly 1.".format( node_name, axis.size) axis = axis.item() diff --git a/tools/mo/unit_tests/mo/utils/ir_reader/ops_test.py b/tools/mo/unit_tests/mo/utils/ir_reader/ops_test.py index aa1f690f4a0..2df88dab8bd 100644 --- a/tools/mo/unit_tests/mo/utils/ir_reader/ops_test.py +++ b/tools/mo/unit_tests/mo/utils/ir_reader/ops_test.py @@ -6,6 +6,7 @@ import tempfile import numpy as np from pathlib import Path +import openvino.runtime.opset12 as opset12 import openvino.runtime.opset11 as opset11 import openvino.runtime.opset10 as opset10 from openvino.runtime import Model, serialize, Core, PartialShape, Dimension @@ -207,3 +208,40 @@ class TestOps(unittest.TestCase): graph = TestOps.check_graph_can_save(model, 'scatter_dynamic_model') scatter_update_node = graph.get_op_nodes(op="ScatterUpdate")[0] self.assertListEqual(scatter_update_node.out_port(0).data.get_value().tolist(), [0, None]) + + def test_pad_12(self): + data_parameter = opset12.parameter([6, 12, 10, 24], name="Data", dtype=np.float32) + pad = opset12.pad(data_parameter, np.int64([0, 0, -1, -2]), np.int64([0, 0, -3, -4]), "constant") + model = Model(pad, [data_parameter]) + graph = TestOps.check_graph_can_save(model, 'pad_model') + pad_node = graph.get_op_nodes(op="Pad")[0] + self.assertEqual(pad_node["version"], "opset12") + self.assertListEqual(pad_node.in_port(1).data.get_value().tolist(), [0, 0, -1, -2]) + self.assertListEqual(pad_node.in_port(2).data.get_value().tolist(), [0, 0, -3, -4]) + self.assertListEqual(pad_node.out_port(0).data.get_shape().tolist(), [6, 12, 6, 18]) + + def test_scatter_elements_update_12(self): + data_parameter = opset12.parameter([10], name="Data", dtype=np.float32) + scatter = opset12.scatter_elements_update(data_parameter, np.int32([5, 0, 7, 5]), np.float32([5., 6., 1.5, -5.]), np.int32(0), "sum", False) + model = Model(scatter, [data_parameter]) + graph = TestOps.check_graph_can_save(model, 'scatter_model') + scatter_node = graph.get_op_nodes(op="ScatterElementsUpdate")[0] + self.assertListEqual(scatter_node.out_port(0).data.get_shape().tolist(), [10]) + self.assertEqual(scatter_node["version"], "opset12") + self.assertEqual(scatter_node['reduction'], 'sum') + self.assertFalse(scatter_node['use_init_val']) + + def test_group_norm_12(self): + data_parameter = opset12.parameter([1, 3, 3, 3], name="Data", dtype=np.float32) + scale = np.array((1, 1, 1), dtype=np.float32) + bias = np.array((1, 1, 1), dtype=np.float32) + num_groups = 1 + epsilon = 1e-6 + node = opset12.group_normalization(data_parameter, scale, bias, num_groups, epsilon) + model = Model(node, [data_parameter]) + graph = TestOps.check_graph_can_save(model, 'group_norm_model') + gn_node = graph.get_op_nodes(op="GroupNormalization")[0] + self.assertListEqual(gn_node.out_port(0).data.get_shape().tolist(), [1, 3, 3, 3]) + self.assertEqual(gn_node["version"], "opset12") + self.assertEqual(gn_node['num_groups'], 1) + self.assertEqual(gn_node['epsilon'], 1e-06)