codekingpro/portable-devtools
115k
1import onnx
2
3from ..quant_utils import TENSOR_NAME_QUANT_SUFFIX, QuantizedValue, QuantizedValueType, attribute_to_kwarg, ms_domain
4from .base_operator import QuantOperatorBase
5
6
7class QLinearPool(QuantOperatorBase):
8 def __init__(self, onnx_quantizer, onnx_node):
9 super().__init__(onnx_quantizer, onnx_node)
10
11 def quantize(self):
12 node = self.node
13
14 # only try to quantize when given quantization parameters for it
15 (
16 data_found,
17 output_scale_name,
18 output_zp_name,
19 _,
20 _,
21 ) = self.quantizer._get_quantization_params(node.output[0])
22
23 # get quantized input tensor names, quantize input if needed
24 (
25 quantized_input_names,
26 input_zero_point_names,
27 input_scale_names,
28 nodes,
29 ) = self.quantizer.quantize_activation(node, [0])
30
31 if not data_found or quantized_input_names is None:
32 return super().quantize()
33
34 # Create an entry for output quantized value.
35 qlinear_output_name = node.output[0] + TENSOR_NAME_QUANT_SUFFIX
36 quantized_output_value = QuantizedValue(
37 node.output[0],
38 qlinear_output_name,
39 output_scale_name,
40 output_zp_name,
41 QuantizedValueType.Input,
42 )
43 self.quantizer.quantized_value_map[node.output[0]] = quantized_output_value
44
45 # Create qlinear pool node for given type (AveragePool, etc)
46 kwargs = {}
47 for attribute in node.attribute:
48 kwargs.update(attribute_to_kwarg(attribute))
49 kwargs["domain"] = ms_domain
50 qlinear_node_name = node.name + "_quant" if node.name else ""
51 qnode = onnx.helper.make_node(
52 "QLinear" + node.op_type,
53 [
54 quantized_input_names[0],
55 input_scale_names[0],
56 input_zero_point_names[0],
57 output_scale_name,
58 output_zp_name,
59 ],
60 [qlinear_output_name],
61 qlinear_node_name,
62 **kwargs,
63 )
64
65 # add all newly created nodes
66 nodes.append(qnode)
67 self.quantizer.new_nodes += nodes
68 