codekingpro/portable-devtools
114k
1import onnx
2
3from ..quant_utils import TENSOR_NAME_QUANT_SUFFIX, QuantizedValue, QuantizedValueType, attribute_to_kwarg, ms_domain
4from .base_operator import QuantOperatorBase
5from .qdq_base_operator import QDQOperatorBase
6
7
8class QLinearWhere(QuantOperatorBase):
9 def should_quantize(self):
10 return True
11
12 def quantize(self):
13 node = self.node
14 assert node.op_type == "Where"
15 if not self.quantizer.force_quantize_no_input_check:
16 self.quantizer.new_nodes += [node]
17 return
18 (
19 data_found,
20 output_scale_name,
21 output_zp_name,
22 _,
23 _,
24 ) = self.quantizer._get_quantization_params(node.output[0])
25 (
26 q_input_names,
27 zero_point_names,
28 scale_names,
29 nodes,
30 ) = self.quantizer.quantize_activation(node, [1, 2])
31 if not data_found or q_input_names is None:
32 return super().quantize()
33 qlinear_output = node.output[0] + TENSOR_NAME_QUANT_SUFFIX
34 qlinear_output_name = node.name + "_quant" if node.name else ""
35
36 q_output = QuantizedValue(
37 node.output[0],
38 qlinear_output,
39 output_scale_name,
40 output_zp_name,
41 QuantizedValueType.Input,
42 )
43 self.quantizer.quantized_value_map[node.output[0]] = q_output
44
45 kwargs = {}
46 for attribute in node.attribute:
47 kwargs.update(attribute_to_kwarg(attribute))
48 kwargs["domain"] = ms_domain
49
50 qlwhere_inputs = [
51 node.input[0],
52 q_input_names[0],
53 scale_names[0],
54 zero_point_names[0],
55 q_input_names[1],
56 scale_names[1],
57 zero_point_names[1],
58 output_scale_name,
59 output_zp_name,
60 ]
61 qlwhere_node = onnx.helper.make_node(
62 "QLinearWhere", qlwhere_inputs, [qlinear_output], qlinear_output_name, **kwargs
63 )
64
65 self.quantizer.new_nodes += nodes
66 self.quantizer.new_nodes += [qlwhere_node]
67
68
69class QDQWhere(QDQOperatorBase):
70 def quantize(self):
71 node = self.node
72 assert node.op_type == "Where"
73 if self.quantizer.force_quantize_no_input_check:
74 if not self.quantizer.is_tensor_quantized(node.input[1]):
75 self.quantizer.quantize_activation_tensor(node.input[1])
76 if not self.quantizer.is_tensor_quantized(node.input[2]):
77 self.quantizer.quantize_activation_tensor(node.input[2])
78 if not self.disable_qdq_for_node_output:
79 for output in node.output:
80 self.quantizer.quantize_activation_tensor(output)
81 elif (
82 self.quantizer.is_tensor_quantized(node.input[1])
83 and self.quantizer.is_tensor_quantized(node.input[2])
84 and not self.disable_qdq_for_node_output
85 ):
86 for output in node.output:
87 self.quantizer.quantize_activation_tensor(output)
88 