Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
where.py88 linesDownload Raw Back to operators
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 
codekingpro/portable-devtools · Team Ai