Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
gemm.py173 linesDownload Raw Back to operators
1import logging
2
3import numpy as np  # noqa: F401
4import onnx
5
6from ..quant_utils import (
7    TENSOR_NAME_QUANT_SUFFIX,
8    QuantizedValue,
9    QuantizedValueType,
10    attribute_to_kwarg,
11    find_by_name,  # noqa: F401
12    get_mul_node,  # noqa: F401
13    ms_domain,
14)
15from .base_operator import QuantOperatorBase  # noqa: F401
16from .matmul import QOpMatMul
17from .qdq_base_operator import QDQOperatorBase
18
19
20def is_B_transposed(gemm_node):  # noqa: N802
21    transB_attribute = [attr for attr in gemm_node.attribute if attr.name == "transB"]  # noqa: N806
22    if transB_attribute:
23        return onnx.helper.get_attribute_value(transB_attribute[0]) > 0
24
25    return False
26
27
28def get_beta(gemm_node):
29    beta_attribute = [attr for attr in gemm_node.attribute if attr.name == "beta"]
30    if beta_attribute:
31        return onnx.helper.get_attribute_value(beta_attribute[0])
32
33    return 1.0
34
35
36def set_default_beta(gemm_node):
37    beta_attribute = [attr for attr in gemm_node.attribute if attr.name == "beta"]
38    if beta_attribute:
39        beta_attribute[0].f = 1.0
40
41    return 1.0
42
43
44class QLinearGemm(QOpMatMul):
45    def __init__(self, onnx_quantizer, onnx_node):
46        super().__init__(onnx_quantizer, onnx_node)
47
48    def quantize(self):
49        node = self.node
50        assert node.op_type == "Gemm"
51
52        (
53            data_found,
54            output_scale_name,
55            output_zp_name,
56            _,
57            _,
58        ) = self.quantizer._get_quantization_params(node.output[0])
59
60        if self.quantizer.is_input_a_initializer(node.input[1]) and self.quantizer.is_per_channel():
61            (
62                quantized_input_names,
63                zero_point_names,
64                scale_names,
65                nodes,
66            ) = self.quantizer.quantize_activation(node, [0])
67            quant_weight_tuple = self.quantizer.quantize_weight_per_channel(
68                node.input[1],
69                self.quantizer.weight_qType,
70                0 if is_B_transposed(node) else 1,
71            )
72            quantized_input_names.append(quant_weight_tuple[0])
73            zero_point_names.append(quant_weight_tuple[1])
74            scale_names.append(quant_weight_tuple[2])
75        else:
76            #  Get Quantized from both activation(input[0]) and weight(input[1])
77            (
78                quantized_input_names,
79                zero_point_names,
80                scale_names,
81                nodes,
82            ) = self.quantizer.quantize_activation(node, [0])
83
84            (
85                quantized_input_names_weight,
86                zero_point_names_weight,
87                scale_names_weight,
88                nodes_weight,
89            ) = self.quantizer.quantize_weight(node, [1], reduce_range=self.quantizer.reduce_range)
90            quantized_input_names.extend(quantized_input_names_weight)
91            zero_point_names.extend(zero_point_names_weight)
92            scale_names.extend(scale_names_weight)
93            nodes.extend(nodes_weight)
94
95        if not data_found or quantized_input_names is None:
96            return super().quantize()
97
98        quantized_bias_name = ""
99        if len(node.input) == 3:
100            if not self.quantizer.is_input_a_initializer(node.input[2]):
101                return super().quantize()
102
103            # Note: if the quantized type is float 8, the bias is converted into float 16.
104            # cublasLtMatMul only supports (b)float16 or float32 bias.
105            quantized_bias_name = self.quantizer.quantize_bias_static(
106                node.input[2], node.input[0], node.input[1], get_beta(self.node)
107            )
108
109        qgemm_output = node.output[0] + TENSOR_NAME_QUANT_SUFFIX
110        qgemm_name = node.name + "_quant" if node.name else ""
111
112        kwargs = {}
113        for attribute in node.attribute:
114            if attribute.name != "beta":
115                kwargs.update(attribute_to_kwarg(attribute))
116        kwargs["domain"] = ms_domain
117
118        # generate input
119        qgemm_inputs = []
120        for i in range(2):
121            qgemm_inputs.extend([quantized_input_names[i], scale_names[i], zero_point_names[i]])
122
123        qgemm_inputs.extend([quantized_bias_name, output_scale_name, output_zp_name])
124
125        qgemm_node = onnx.helper.make_node("QGemm", qgemm_inputs, [qgemm_output], qgemm_name, **kwargs)
126        nodes.append(qgemm_node)
127
128        # Create an entry for this quantized value
129        q_output = QuantizedValue(
130            node.output[0],
131            qgemm_output,
132            output_scale_name,
133            output_zp_name,
134            QuantizedValueType.Input,
135            node_type=node.op_type,
136            node_qtype=self.quantizer.weight_qType,
137        )
138        self.quantizer.quantized_value_map[node.output[0]] = q_output
139
140        self.quantizer.new_nodes += nodes
141
142
143class QDQGemm(QDQOperatorBase):
144    def __init__(self, onnx_quantizer, onnx_node):
145        super().__init__(onnx_quantizer, onnx_node)
146
147    def quantize(self):
148        node = self.node
149        assert node.op_type == "Gemm"
150
151        self.quantizer.quantize_activation_tensor(node.input[0])
152        if not self.disable_qdq_for_node_output:
153            self.quantizer.quantize_activation_tensor(node.output[0])
154
155        is_weight_per_channel, weight_axis = self.quantizer.is_tensor_per_channel(
156            node.input[1], default_axis=0 if is_B_transposed(node) else 1
157        )
158        if is_weight_per_channel:
159            self.quantizer.quantize_weight_tensor_per_channel(node.input[1], weight_axis)
160        else:
161            self.quantizer.quantize_weight_tensor(node.input[1])
162
163        if len(node.input) == 3:
164            if self.quantizer.is_input_a_initializer(node.input[2]):
165                self.quantizer.quantize_bias_tensor(
166                    node.name, node.input[2], node.input[0], node.input[1], get_beta(self.node)
167                )
168                set_default_beta(self.node)
169            else:
170                logging.warning(
171                    f"Bias of Gemm node '{self.node.name}' is not constant. Please exclude this node for better performance."
172                )
173 
codekingpro/portable-devtools · Team Ai