Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes15kdownloads
split.py64 linesDownload Raw Back to operators
1import onnx
2
3from ..quant_utils import QuantizedValue, QuantizedValueType, attribute_to_kwarg
4from .base_operator import QuantOperatorBase
5from .qdq_base_operator import QDQOperatorBase
6
7
8class QSplit(QuantOperatorBase):
9    def __init__(self, onnx_quantizer, onnx_node):
10        super().__init__(onnx_quantizer, onnx_node)
11
12    def quantize(self):
13        node = self.node
14        (
15            quantized_input_names,
16            zero_point_names,
17            scale_names,
18            nodes,
19        ) = self.quantizer.quantize_activation(node, [0])
20        if quantized_input_names is None:
21            return super().quantize()
22
23        quantized_node_name = ""
24        if node.name:
25            quantized_node_name = node.name + "_quant"
26        kwargs = {}
27        for attribute in node.attribute:
28            kwargs.update(attribute_to_kwarg(attribute))
29
30        # Output just derive the scale/zero from input
31        quantized_output_names = []
32        for output_name in node.output:
33            quantized_output_name = output_name + "quantized"
34            quantized_output_names.append(quantized_output_name)
35            q_output = QuantizedValue(
36                output_name,
37                quantized_output_name,
38                scale_names[0],
39                zero_point_names[0],
40                QuantizedValueType.Input,
41            )
42            self.quantizer.quantized_value_map[output_name] = q_output
43
44        if len(node.input) > 1:
45            quantized_input_names.extend(node.input[1:])
46        quantized_node = onnx.helper.make_node(
47            node.op_type, quantized_input_names, quantized_output_names, quantized_node_name, **kwargs
48        )
49
50        nodes.append(quantized_node)
51        self.quantizer.new_nodes += nodes
52
53
54class QDQSplit(QDQOperatorBase):
55    def quantize(self):
56        node = self.node
57        assert node.op_type == "Split"
58
59        if not self.quantizer.is_tensor_quantized(node.input[0]):
60            self.quantizer.quantize_activation_tensor(node.input[0])
61        if not self.disable_qdq_for_node_output:
62            for output in node.output:
63                self.quantizer.quantize_output_same_as_input(output, node.input[0], node.name)
64 
codekingpro/portable-devtools · Team Ai