Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
lstm.py122 linesDownload Raw Back to operators
1import numpy
2import onnx
3from onnx import onnx_pb as onnx_proto
4
5from ..quant_utils import QuantType, attribute_to_kwarg, ms_domain  # noqa: F401
6from .base_operator import QuantOperatorBase
7
8"""
9    Quantize LSTM
10"""
11
12
13class LSTMQuant(QuantOperatorBase):
14    def __init__(self, onnx_quantizer, onnx_node):
15        super().__init__(onnx_quantizer, onnx_node)
16
17    def quantize(self):
18        """
19        parameter node: LSTM node.
20        parameter new_nodes_list: List of new nodes created before processing this node.
21        return: a list of nodes in topological order that represents quantized Attention node.
22        """
23        node = self.node
24        assert node.op_type == "LSTM"
25
26        if not self.quantizer.is_valid_quantize_weight(node.input[1]) or not self.quantizer.is_valid_quantize_weight(
27            node.input[2]
28        ):
29            super().quantize()
30            return
31
32        model = self.quantizer.model
33        W = model.get_initializer(node.input[1])  # noqa: N806
34        R = model.get_initializer(node.input[2])  # noqa: N806
35
36        if len(W.dims) != 3 or len(R.dims) != 3:
37            super().quantize()
38            return
39
40        [W_num_dir, W_4_hidden_size, W_input_size] = W.dims  # noqa: N806
41        [R_num_dir, R_4_hidden_size, R_hidden_size] = R.dims  # noqa: N806
42
43        if self.quantizer.is_per_channel():
44            del W.dims[0]
45            del R.dims[0]
46            W.dims[0] = W_num_dir * W_4_hidden_size
47            R.dims[0] = R_num_dir * R_4_hidden_size
48
49        quant_input_weight_tuple = self.quantizer.quantize_weight_per_channel(
50            node.input[1],
51            onnx_proto.TensorProto.INT8,
52            0,  # self.quantizer.weight_qType?
53        )
54        quant_recurrent_weight_tuple = self.quantizer.quantize_weight_per_channel(
55            node.input[2],
56            onnx_proto.TensorProto.INT8,
57            0,  # self.quantizer.weight_qType?
58        )
59
60        W_quant_weight = model.get_initializer(quant_input_weight_tuple[0])  # noqa: N806
61        R_quant_weight = model.get_initializer(quant_recurrent_weight_tuple[0])  # noqa: N806
62
63        W_quant_array = onnx.numpy_helper.to_array(W_quant_weight)  # noqa: N806
64        R_quant_array = onnx.numpy_helper.to_array(R_quant_weight)  # noqa: N806
65
66        W_quant_array = numpy.reshape(W_quant_array, (W_num_dir, W_4_hidden_size, W_input_size))  # noqa: N806
67        R_quant_array = numpy.reshape(R_quant_array, (R_num_dir, R_4_hidden_size, R_hidden_size))  # noqa: N806
68
69        W_quant_array = numpy.transpose(W_quant_array, (0, 2, 1))  # noqa: N806
70        R_quant_array = numpy.transpose(R_quant_array, (0, 2, 1))  # noqa: N806
71
72        W_quant_tranposed = onnx.numpy_helper.from_array(W_quant_array, quant_input_weight_tuple[0])  # noqa: N806
73        R_quant_tranposed = onnx.numpy_helper.from_array(R_quant_array, quant_recurrent_weight_tuple[0])  # noqa: N806
74
75        model.remove_initializers([W_quant_weight, R_quant_weight])
76        model.add_initializer(W_quant_tranposed)
77        model.add_initializer(R_quant_tranposed)
78
79        W_quant_zp = model.get_initializer(quant_input_weight_tuple[1])  # noqa: N806
80        R_quant_zp = model.get_initializer(quant_recurrent_weight_tuple[1])  # noqa: N806
81        W_quant_scale = model.get_initializer(quant_input_weight_tuple[2])  # noqa: N806
82        R_quant_scale = model.get_initializer(quant_recurrent_weight_tuple[2])  # noqa: N806
83
84        if self.quantizer.is_per_channel():
85            W_quant_zp.dims[:] = [W_num_dir, W_4_hidden_size]
86            R_quant_zp.dims[:] = [R_num_dir, R_4_hidden_size]
87            W_quant_scale.dims[:] = [W_num_dir, W_4_hidden_size]
88            R_quant_scale.dims[:] = [R_num_dir, R_4_hidden_size]
89
90        inputs = []
91        input_len = len(node.input)
92        inputs.extend([node.input[0]])
93        inputs.extend([quant_input_weight_tuple[0], quant_recurrent_weight_tuple[0]])
94        inputs.extend([node.input[3] if input_len > 3 else ""])
95        inputs.extend([node.input[4] if input_len > 4 else ""])
96        inputs.extend([node.input[5] if input_len > 5 else ""])
97        inputs.extend([node.input[6] if input_len > 6 else ""])
98        inputs.extend([node.input[7] if input_len > 7 else ""])
99        inputs.extend(
100            [
101                quant_input_weight_tuple[2],
102                quant_input_weight_tuple[1],
103                quant_recurrent_weight_tuple[2],
104                quant_recurrent_weight_tuple[1],
105            ]
106        )
107
108        kwargs = {}
109        for attribute in node.attribute:
110            if attribute.name == "layout":
111                continue
112            kwargs.update(attribute_to_kwarg(attribute))
113        kwargs["domain"] = ms_domain
114
115        quant_lstm_name = "" if not node.name else node.name + "_quant"
116        quant_lstm_node = onnx.helper.make_node("DynamicQuantizeLSTM", inputs, node.output, quant_lstm_name, **kwargs)
117        self.quantizer.new_nodes.append(quant_lstm_node)
118
119        dequantize_node = self.quantizer._dequantize_value(node.input[0])
120        if dequantize_node is not None:
121            self.quantizer.new_nodes.append(dequantize_node)
122 
codekingpro/portable-devtools · Team Ai