codekingpro/portable-devtools
114k
1from ..quant_utils import TENSOR_NAME_QUANT_SUFFIX, QuantizedValue, QuantizedValueType
2from .base_operator import QuantOperatorBase
3from .qdq_base_operator import QDQOperatorBase
4
5"""
6 Quantize Gather
7"""
8
9
10class GatherQuant(QuantOperatorBase):
11 def __init__(self, onnx_quantizer, onnx_node):
12 super().__init__(onnx_quantizer, onnx_node)
13
14 def should_quantize(self):
15 if not self.quantizer.should_quantize_node(self.node):
16 return False
17
18 return self.quantizer.is_valid_quantize_weight(self.node.input[0])
19
20 def quantize(self):
21 node = self.node
22 assert node.op_type == "Gather"
23
24 (
25 quantized_input_names,
26 zero_point_names,
27 scale_names,
28 nodes,
29 ) = self.quantizer.quantize_activation(node, [0])
30 if quantized_input_names is None:
31 return super().quantize()
32
33 gather_new_output = node.output[0] + TENSOR_NAME_QUANT_SUFFIX
34
35 # Create an entry for this quantized value
36 q_output = QuantizedValue(
37 node.output[0],
38 gather_new_output,
39 scale_names[0],
40 zero_point_names[0],
41 QuantizedValueType.Input,
42 )
43 self.quantizer.quantized_value_map[node.output[0]] = q_output
44
45 node.output[0] = gather_new_output
46 node.input[0] = quantized_input_names[0]
47 nodes.append(node)
48
49 self.quantizer.new_nodes += nodes
50
51
52class QDQGather(QDQOperatorBase):
53 def __init__(self, onnx_quantizer, onnx_node):
54 super().__init__(onnx_quantizer, onnx_node)
55
56 def quantize(self):
57 node = self.node
58 assert node.op_type == "Gather" or node.op_type == "GatherElements"
59
60 if self.quantizer.is_valid_quantize_weight(node.input[0]) or self.quantizer.force_quantize_no_input_check:
61 self.quantizer.quantize_activation_tensor(node.input[0])
62 self.quantizer.quantize_output_same_as_input(node.output[0], node.input[0], node.name)
63 elif self.quantizer.is_tensor_quantized(node.input[0]):
64 self.quantizer.quantize_output_same_as_input(node.output[0], node.input[0], node.name)
65 