codekingpro/portable-devtools
114k
1import logging
2
3import onnx
4from onnx import onnx_pb as onnx_proto # noqa: F401
5
6from ..quant_utils import attribute_to_kwarg, ms_domain
7from .base_operator import QuantOperatorBase
8
9"""
10Quantizes the EmbedLayerNorm fused ONNXRuntime Op.
11
12This Quant operator keeps the input and segment IDs at int32 but will quantize all initializer and
13weight inputs associated with the node to uint8.
14"""
15
16
17class EmbedLayerNormalizationQuant(QuantOperatorBase):
18 def __init__(self, onnx_quantizer, onnx_node):
19 super().__init__(onnx_quantizer, onnx_node)
20
21 def should_quantize(self):
22 return self.quantizer.should_quantize_node(self.node)
23
24 def quantize(self):
25 node = self.node
26 assert node.op_type == "EmbedLayerNormalization"
27
28 if len(node.output) > 2:
29 logging.info(f"Quantization is not applied to {node.name} since it has 3 outputs")
30 return super().quantize()
31
32 """
33 Pre-quantization EmbedLayerNorm inputs:
34 [0] input_ids (int32)
35 [1] segment_ids (int32)
36 [2] word_embedding (float32)
37 [3] position_embedding (float32)
38 [4] segment_embedding (float32)
39 [5] gamma (float32)
40 [6] beta (float32)
41 [7] mask (int32) (optional)
42 """
43 (
44 quantized_input_names,
45 zero_point_names,
46 scale_names,
47 nodes,
48 ) = self.quantizer.quantize_activation(node, [2, 3, 4, 5, 6])
49 if quantized_input_names is None:
50 return super().quantize()
51
52 qembed_layer_norm_name = "" if not node.name else node.name + "_quant"
53
54 """
55 Quantized Input Tensor List
56 [0] input_ids (int32)
57 [1] segment_ids (int32)
58 [2] word_embedding (uint8)
59 [3] position_embedding (uint8)
60 [4] segment_embedding (uint8)
61 [5] gamma (uint8)
62 [6] beta (uint8)
63 [7] mask (int32) (optional)
64 [8] word_embedding_scale (float)
65 [9] position_embedding_scale (float)
66 [10] segment_embedding_scale (float)
67 [11] gamma_scale (float)
68 [12] beta_scale (float)
69 [13] word_embedding_zero_point (uint8)
70 [14] position_embedding_zero_point (uint8)
71 [15] segment_embedding_zero_point (uint8)
72 [16] gamma_zero_point (uint8)
73 [17] beta_zero_point (uint8)
74 """
75 inputs = []
76 # 'input_ids'
77 inputs.extend([node.input[0]])
78 # 'segment_ids'
79 inputs.extend([node.input[1]])
80 # 'word_embedding_quant'
81 inputs.extend([quantized_input_names[0]])
82 # 'position_embedding_quant'
83 inputs.extend([quantized_input_names[1]])
84 # 'segment_embedding_quant'
85 inputs.extend([quantized_input_names[2]])
86 # 'gamma_quant'
87 inputs.extend([quantized_input_names[3]])
88 # 'beta_quant'
89 inputs.extend([quantized_input_names[4]])
90 # 'mask' (optional)
91 inputs.extend([node.input[7] if len(node.input) > 7 else ""])
92
93 # Add all scales:
94 inputs.extend([scale_names[0]])
95 inputs.extend([scale_names[1]])
96 inputs.extend([scale_names[2]])
97 inputs.extend([scale_names[3]])
98 inputs.extend([scale_names[4]])
99
100 # Add all zero points:
101 inputs.extend([zero_point_names[0]])
102 inputs.extend([zero_point_names[1]])
103 inputs.extend([zero_point_names[2]])
104 inputs.extend([zero_point_names[3]])
105 inputs.extend([zero_point_names[4]])
106
107 kwargs = {}
108 for attribute in node.attribute:
109 kwargs.update(attribute_to_kwarg(attribute))
110 kwargs["domain"] = ms_domain
111
112 qembed_layer_norm_node = onnx.helper.make_node(
113 "QEmbedLayerNormalization",
114 inputs,
115 node.output,
116 qembed_layer_norm_name,
117 **kwargs,
118 )
119 nodes.append(qembed_layer_norm_node)
120
121 self.quantizer.new_nodes += nodes
122 