Add opset12 operations to MO IR Reader (#18851)

This commit is contained in:
Maxim Vafin 2023-07-28 19:14:38 +02:00 committed by GitHub
parent 6a94ae3409
commit e3f19b59e7
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
3 changed files with 73 additions and 1 deletions

View File

@ -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']

View File

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

View File

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