codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6from .qdq_base_operator import QDQOperatorBase
7
8
9class QDQNormalization(QDQOperatorBase):
10 def __init__(self, onnx_quantizer, onnx_node):
11 super().__init__(onnx_quantizer, onnx_node)
12
13 def quantize(self):
14 node = self.node
15 assert node.op_type in {"InstanceNormalization", "LayerNormalization", "BatchNormalization"}
16
17 # Input
18 self.quantizer.quantize_activation_tensor(node.input[0])
19
20 # Scale
21 scale_is_initializer = self.quantizer.is_input_a_initializer(node.input[1])
22 scale_is_per_channel, scale_channel_axis = self.quantizer.is_tensor_per_channel(
23 node.input[1], default_axis=1, op_type=node.op_type
24 )
25
26 if scale_is_per_channel:
27 self.quantizer.quantize_weight_tensor_per_channel(node.input[1], axis=scale_channel_axis)
28 elif scale_is_initializer:
29 self.quantizer.quantize_weight_tensor(node.input[1])
30 else:
31 self.quantizer.quantize_activation_tensor(node.input[1])
32
33 # Bias
34 if len(node.input) > 2 and node.input[2]:
35 self.quantizer.quantize_bias_tensor(node.name, node.input[2], node.input[0], node.input[1])
36
37 # Output
38 if not self.disable_qdq_for_node_output:
39 for output_name in node.output:
40 self.quantizer.quantize_activation_tensor(output_name)
41 