codekingpro/portable-devtools
114k
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 