[PT FE] Add quantized::cat (#18750)

* Add quantized cat

* Add quantized cat

* Format

* Improve tests

* Fix merge conflict
This commit is contained in:
Mateusz Mikolajczyk 2023-07-27 10:32:17 +02:00 committed by GitHub
parent f8f6d4a5b5
commit a0a1102499
No known key found for this signature in database
GPG Key ID: 4AEE18F83AFDEB23
4 changed files with 180 additions and 15 deletions

View File

@ -7,6 +7,7 @@
#include "openvino/op/parameter.hpp"
#include "pt_framework_node.hpp"
#include "utils.hpp"
#include "utils_quantize.hpp"
namespace ov {
namespace frontend {
@ -22,15 +23,17 @@ OutputVector translate_cat_common(const NodeContext& context,
// couldn't get list elements
auto fw_node = std::make_shared<PtFrameworkNode>(context.get_decoder(), OutputVector{context.get_input(0)}, 1);
auto attrs = fw_node->get_attrs();
// If this fails it means axis is dynamic and aten::cat will be converted to fw node in regular pipeline
// If this fails it means axis is dynamic and <aten/quantized>::cat will be converted to fw node in regular
// pipeline
attrs["axis"] = std::to_string(axis);
fw_node->set_attrs(attrs);
return {context.mark_node(fw_node)};
} else {
auto first_elem = list_elems.front().get_node_shared_ptr();
FRONT_END_OP_CONVERSION_CHECK(list_elems.size() > 1 || !ov::as_type_ptr<v0::Parameter>(first_elem),
"aten::cat is located inside body while inputs are located outside of the body. "
"This case is not supported.");
FRONT_END_OP_CONVERSION_CHECK(
list_elems.size() > 1 || !ov::as_type_ptr<v0::Parameter>(first_elem),
"<aten/quantized>::cat is located inside body while inputs are located outside of the body. "
"This case is not supported.");
}
auto concat = std::make_shared<v0::Concat>(OutputVector(list_elems.begin(), list_elems.end()), axis);
return {context.mark_node(concat)};
@ -55,6 +58,18 @@ OutputVector translate_cat_fx(const NodeContext& context) {
return translate_cat_common(context, list_elems, axis);
};
OutputVector translate_quantized_cat(const NodeContext& context) {
num_inputs_check(context, 4, 4);
const auto&& list_elems = get_list_as_outputs(context.get_input(0));
auto axis = context.const_input<int64_t>(1);
FRONT_END_OP_CONVERSION_CHECK(!list_elems.empty(), "Couldn't find quantized input for quantized::cat operation.");
return {quantize(context,
translate_cat_common(context, list_elems, axis)[0],
context.get_input(2),
context.get_input(3),
list_elems.front())};
};
} // namespace op
} // namespace pytorch
} // namespace frontend

View File

@ -172,6 +172,7 @@ OP_CONVERTER(translate_var_mean);
OP_CONVERTER(translate_where);
OP_CONVERTER(translate_zeros);
OP_CONVERTER(translate_zeros_like);
OP_CONVERTER(translate_quantized_cat);
OP_CONVERTER(translate_quantized_convnd);
OP_CONVERTER(translate_quantized_convnd_relu);
OP_CONVERTER(translate_quantized_linear);
@ -455,6 +456,7 @@ const std::map<std::string, CreatorFunction> get_supported_ops_ts() {
{"prim::type", op::skip_node}, // Used with prim::device, pass PtFrameworkNode.
{"quantized::add", op::translate_quantized_add},
{"quantized::add_relu", op::translate_quantized_add_relu},
{"quantized::cat", op::translate_quantized_cat},
{"quantized::conv2d", op::translate_quantized_convnd},
{"quantized::conv2d_relu", op::translate_quantized_convnd_relu},
{"quantized::hardswish", op::translate_quantized_hardswish},

View File

@ -37,6 +37,8 @@ AtenCatToConcat::AtenCatToConcat() {
ov::matcher_pass_callback callback = [](ov::pass::pattern::Matcher& m) {
auto cat = cast_fw_node(m.get_match_root(), "aten::cat");
if (!cat)
cat = cast_fw_node(m.get_match_root(), "quantized::cat");
if (!cat)
return false;
@ -45,19 +47,19 @@ AtenCatToConcat::AtenCatToConcat() {
auto axis_node = cat->get_input_node_shared_ptr(1);
auto axis_const = std::dynamic_pointer_cast<v0::Constant>(axis_node);
if (!axis_const) {
add_exception_to_fw_node(cat, "aten::cat unsupported case: axis is not a constant.");
add_exception_to_fw_node(cat, "<aten/quantized>::cat unsupported case: axis is not a constant.");
return false;
}
auto _axis = axis_const->cast_vector<int64_t>();
if (_axis.size() != 1) {
add_exception_to_fw_node(cat, "aten::cat unsupported case: axis is not a scalar.");
add_exception_to_fw_node(cat, "<aten/quantized>::cat unsupported case: axis is not a scalar.");
return false;
}
axis = _axis[0];
} else {
const auto& attrs = cat->get_attrs();
if (attrs.find("axis") == attrs.end()) {
add_exception_to_fw_node(cat, "aten::cat unsupported case: axis not found in attributes.");
add_exception_to_fw_node(cat, "<aten/quantized>::cat unsupported case: axis not found in attributes.");
return false;
}
axis = std::stoll(attrs.at("axis"));
@ -81,14 +83,15 @@ AtenCatToConcat::AtenCatToConcat() {
if (!append) {
add_exception_to_fw_node(
cat,
"aten::cat unsupported case: aten::append wasn't found inside prim::Loop body.");
"<aten/quantized>::cat unsupported case: aten::append wasn't found inside prim::Loop body.");
return false;
}
auto param = std::dynamic_pointer_cast<v0::Parameter>(append->get_input_node_shared_ptr(0));
if (!param) {
add_exception_to_fw_node(cat,
"aten::cat unsupported case: input of aten::append inside prim::Loop "
"body is not a body input.");
add_exception_to_fw_node(
cat,
"<aten/quantized>::cat unsupported case: input of aten::append inside prim::Loop "
"body is not a body input.");
return false;
}
auto body_param_index = body->get_parameter_index(param);
@ -103,9 +106,10 @@ AtenCatToConcat::AtenCatToConcat() {
FRONT_END_GENERAL_CHECK(input_index >= 0, "Couldn't find descriptor for input.");
auto list_construct = cast_fw_node(loop->get_input_node_shared_ptr(input_index), "prim::ListConstruct");
if (!list_construct || list_construct->get_input_size() > 0) {
add_exception_to_fw_node(cat,
"aten::cat unsupported case: aten::append input outside of prim::Loop "
"body is not a prim::ListConstruct.");
add_exception_to_fw_node(
cat,
"<aten/quantized>::cat unsupported case: aten::append input outside of prim::Loop "
"body is not a prim::ListConstruct.");
return false;
}
auto new_result = std::make_shared<v0::Result>(append->input_value(1));
@ -131,4 +135,4 @@ AtenCatToConcat::AtenCatToConcat() {
} // namespace pass
} // namespace pytorch
} // namespace frontend
} // namespace ov
} // namespace ov

View File

@ -0,0 +1,144 @@
# Copyright (C) 2018-2023 Intel Corporation
# SPDX-License-Identifier: Apache-2.0
import numpy as np
import pytest
import torch
from pytorch_layer_test_class import PytorchLayerTest
class aten_quantized_cat(torch.nn.Module):
def __init__(self, scale, zero_point, dtype):
super().__init__()
self.scale = float(scale)
self.zero_point = int(zero_point)
self.dtype = dtype
def forward(self, inp):
x = torch.quantize_per_tensor(inp, 1.3, 0, self.dtype)
y = torch.quantize_per_tensor(inp, 1.0, 1, self.dtype)
return torch.dequantize(torch.ops.quantized.cat([x, y], 1, self.scale, self.zero_point))
class aten_append_quantized_cat(torch.nn.Module):
def __init__(self, scale, zero_point, dtype):
super().__init__()
self.scale = float(scale)
self.zero_point = int(zero_point)
self.dtype = dtype
def forward(self, x):
x = torch.quantize_per_tensor(x, 1.0, 0, self.dtype)
list = []
list.append(x)
list.append(x)
return torch.dequantize(torch.ops.quantized.cat(list, 1, self.scale, self.zero_point))
class aten_loop_append_quantized_cat(torch.nn.Module):
def __init__(self, scale, zero_point, dtype):
super().__init__()
self.scale = float(scale)
self.zero_point = int(zero_point)
self.dtype = dtype
def forward(self, x):
x = torch.quantize_per_tensor(x, 1.0, 0, self.dtype)
list = []
for i in range(3):
list.append(x)
return torch.dequantize(torch.ops.quantized.cat(list, 1, self.scale, self.zero_point))
class aten_add_quantized_cat(torch.nn.Module):
def __init__(self, scale, zero_point, dtype):
super().__init__()
self.scale = float(scale)
self.zero_point = int(zero_point)
self.dtype = dtype
def forward(self, x):
x = torch.quantize_per_tensor(x, 1.0, 0, self.dtype)
list = [x, x]
list2 = list + [x, x]
return torch.dequantize(torch.ops.quantized.cat(list2, 1, self.scale, self.zero_point))
class TestQuantizedCat(PytorchLayerTest):
def _prepare_input(self):
return (np.random.rand(2, 1, 3).astype(np.float32),)
@pytest.mark.parametrize("scale", [1.0, 0.3, 1.3])
@pytest.mark.parametrize("zero_point", [0, 1])
@pytest.mark.parametrize("dtype", [torch.quint8, torch.qint8])
@pytest.mark.nightly
@pytest.mark.precommit
def test_quantized_cat(self, scale, zero_point, dtype, ie_device, precision, ir_version):
self._test(
aten_quantized_cat(scale, zero_point, dtype),
None,
["quantized::cat", "prim::ListConstruct"],
ie_device,
precision,
ir_version,
quantized_ops=True,
freeze_model=False,
quant_size=scale,
)
@pytest.mark.parametrize("scale", [1, 0.3, 1.3])
@pytest.mark.parametrize("zero_point", [0, 1])
@pytest.mark.parametrize("dtype", [torch.quint8, torch.qint8])
@pytest.mark.nightly
@pytest.mark.precommit
def test_append_quantized_cat(self, scale, zero_point, dtype, ie_device, precision, ir_version):
self._test(
aten_append_quantized_cat(scale, zero_point, dtype),
None,
["quantized::cat", "aten::append", "prim::ListConstruct"],
ie_device,
precision,
ir_version,
quantized_ops=True,
freeze_model=False,
quant_size=scale,
)
@pytest.mark.parametrize("scale", [1, 0.3, 1.3])
@pytest.mark.parametrize("zero_point", [0, 1])
@pytest.mark.parametrize("dtype", [torch.quint8, torch.qint8])
@pytest.mark.nightly
@pytest.mark.precommit
@pytest.mark.xfail(
reason="Transformation RemoveMultiSubGraphOpDanglingParamsResults doesn't support removing unused merged inputs, ticket 112833."
)
def test_loop_append_quantized_cat(self, scale, zero_point, dtype, ie_device, precision, ir_version):
self._test(
aten_loop_append_quantized_cat(scale, zero_point, dtype),
None,
["quantized::cat", "aten::append", "prim::ListConstruct", "prim::Loop"],
ie_device,
precision,
ir_version,
quantized_ops=True,
freeze_model=False,
quant_size=scale,
)
@pytest.mark.parametrize("scale", [1, 0.3, 1.3])
@pytest.mark.parametrize("zero_point", [0, 1])
@pytest.mark.parametrize("dtype", [torch.quint8, torch.qint8])
@pytest.mark.nightly
@pytest.mark.precommit
def test_add_quantized_cat(self, scale, zero_point, dtype, ie_device, precision, ir_version):
self._test(
aten_add_quantized_cat(scale, zero_point, dtype),
None,
["quantized::cat", "aten::add", "prim::ListConstruct"],
ie_device,
precision,
ir_version,
quantized_ops=True,
freeze_model=False,
quant_size=scale,
)