Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
embed_layernorm.py122 linesDownload Raw Back to operators
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 
codekingpro/portable-devtools · Team Ai