Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_model_bert.py513 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6from logging import getLogger
7
8import numpy as np
9from convert_to_packing_mode import PackingMode
10from fusion_attention import AttentionMask, FusionAttention
11from fusion_bart_attention import FusionBartAttention
12from fusion_biasgelu import FusionBiasGelu
13from fusion_constant_fold import FusionConstantFold
14from fusion_embedlayer import FusionEmbedLayerNormalization
15from fusion_fastgelu import FusionFastGelu
16from fusion_gelu import FusionGelu
17from fusion_gelu_approximation import FusionGeluApproximation
18from fusion_gemmfastgelu import FusionGemmFastGelu
19from fusion_layernorm import FusionLayerNormalization, FusionLayerNormalizationTF
20from fusion_options import AttentionMaskFormat, FusionOptions
21from fusion_qordered_attention import FusionQOrderedAttention
22from fusion_qordered_gelu import FusionQOrderedGelu
23from fusion_qordered_layernorm import FusionQOrderedLayerNormalization
24from fusion_qordered_matmul import FusionQOrderedMatMul
25from fusion_quickgelu import FusionQuickGelu
26from fusion_reshape import FusionReshape
27from fusion_rotary_attention import FusionRotaryEmbeddings
28from fusion_shape import FusionShape
29from fusion_simplified_layernorm import FusionSimplifiedLayerNormalization, FusionSkipSimplifiedLayerNormalization
30from fusion_skiplayernorm import FusionBiasSkipLayerNormalization, FusionSkipLayerNormalization
31from fusion_utils import FusionUtils
32from onnx import ModelProto, TensorProto, helper, numpy_helper
33from onnx_model import OnnxModel
34
35logger = getLogger(__name__)
36
37
38class BertOnnxModel(OnnxModel):
39    def __init__(self, model: ModelProto, num_heads: int = 0, hidden_size: int = 0):
40        """Initialize BERT ONNX Model.
41
42        Args:
43            model (ModelProto): the ONNX model
44            num_heads (int, optional): number of attention heads. Defaults to 0 (detect the parameter automatically).
45            hidden_size (int, optional): hidden dimension. Defaults to 0 (detect the parameter automatically).
46        """
47        assert (num_heads == 0 and hidden_size == 0) or (num_heads > 0 and hidden_size % num_heads == 0)
48
49        super().__init__(model)
50        self.num_heads = num_heads
51        self.hidden_size = hidden_size
52
53        self.attention_mask = AttentionMask(self)
54        self.attention_fusion = FusionAttention(self, self.hidden_size, self.num_heads, self.attention_mask)
55        self.qordered_attention_fusion = FusionQOrderedAttention(
56            self, self.hidden_size, self.num_heads, self.attention_mask
57        )
58        self.utils = FusionUtils(self)
59
60    def fuse_constant_fold(self):
61        fusion = FusionConstantFold(self)
62        fusion.apply()
63
64    def fuse_attention(self):
65        self.attention_fusion.apply()
66        # Only relevant in models with Q-DQ nodes
67        self.qordered_attention_fusion.apply()
68
69    def fuse_gelu(self):
70        fusion = FusionGelu(self)
71        fusion.apply()
72        fusion = FusionFastGelu(self)
73        fusion.apply()
74        fusion = FusionQuickGelu(self)
75        fusion.apply()
76        # Only relevant in models with Q-DQ nodes
77        fusion = FusionQOrderedGelu(self)
78        fusion.apply()
79
80    def fuse_bias_gelu(self, is_fastgelu):
81        fusion = FusionBiasGelu(self, is_fastgelu)
82        fusion.apply()
83
84    def gelu_approximation(self):
85        fusion = FusionGeluApproximation(self)
86        fusion.apply()
87
88    def fuse_gemm_fast_gelu(self):
89        fusion = FusionGemmFastGelu(self)
90        fusion.apply()
91
92    def fuse_add_bias_skip_layer_norm(self):
93        fusion = FusionBiasSkipLayerNormalization(self)
94        fusion.apply()
95
96    def fuse_reshape(self):
97        fusion = FusionReshape(self)
98        fusion.apply()
99
100    def fuse_shape(self):
101        fusion = FusionShape(self)
102        fusion.apply()
103
104    def fuse_embed_layer(self, use_mask_index):
105        fusion = FusionEmbedLayerNormalization(self, use_mask_index)
106        fusion.apply()
107
108    def fuse_layer_norm(self):
109        fusion = FusionLayerNormalization(self)
110        fusion.apply()
111
112        fusion = FusionLayerNormalizationTF(self)
113        fusion.apply()
114
115        # Only relevant in models with Q-DQ nodes
116        fusion = FusionQOrderedLayerNormalization(self)
117        fusion.apply()
118
119    def fuse_simplified_layer_norm(self):
120        fusion = FusionSimplifiedLayerNormalization(self)
121        fusion.apply()
122
123    def fuse_skip_layer_norm(self, shape_infer=True):
124        fusion = FusionSkipLayerNormalization(self, shape_infer=shape_infer)
125        fusion.apply()
126
127    def fuse_skip_simplified_layer_norm(self):
128        fusion = FusionSkipSimplifiedLayerNormalization(self)
129        fusion.apply()
130
131    def fuse_rotary_embeddings(self):
132        fusion = FusionRotaryEmbeddings(self)
133        fusion.apply()
134        # Remove non-MS domain functions
135        rot_emb_nodes = list(
136            filter(
137                lambda node: node.op_type == "RotaryEmbedding" and node.domain != "com.microsoft",
138                self.model.graph.node,
139            )
140        )
141        non_ms_domains_to_keep = {node.domain for node in rot_emb_nodes}
142        i = 0
143        while i < len(self.model.functions):
144            fn = self.model.functions[i]
145            if "RotaryEmbedding" in fn.name and fn.domain not in non_ms_domains_to_keep:
146                self.model.functions.remove(fn)
147            else:
148                i += 1
149
150    # Only relevant in models with Q-DQ nodes
151    def fuse_qordered_mamtul(self):
152        fusion = FusionQOrderedMatMul(self)
153        fusion.apply()
154
155    def get_graph_inputs_from_node_type(self, op_type: str, input_indices: list[int], casted: bool):
156        """
157        Get graph inputs that feed into node type (like EmbedLayerNormalization or Attention).
158        Returns a list of the graph input names based on the filter whether it is casted or not.
159        """
160        graph_inputs = []
161
162        output_name_to_node = self.output_name_to_node()
163        nodes = self.get_nodes_by_op_type(op_type)
164        for node in nodes:
165            bert_inputs = [node.input[i] for i in input_indices if i < len(node.input)]
166            for bert_input in bert_inputs:
167                if self.find_graph_input(bert_input):
168                    if not casted:
169                        graph_inputs.append(bert_input)
170                elif bert_input in output_name_to_node:
171                    parent = output_name_to_node[bert_input]
172                    if parent.op_type == "Cast" and self.find_graph_input(parent.input[0]) is not None:
173                        if casted:
174                            graph_inputs.append(parent.input[0])
175        return graph_inputs
176
177    def get_graph_inputs_from_fused_nodes(self, casted: bool):
178        inputs = self.get_graph_inputs_from_node_type("EmbedLayerNormalization", [0, 1, 7], casted)
179        inputs += self.get_graph_inputs_from_node_type("Attention", [3], casted)
180        return inputs
181
182    def change_graph_inputs_to_int32(self):
183        """Change data type of all graph inputs to int32 type, and add Cast node if needed."""
184        graph = self.graph()
185        add_cast_count = 0
186        remove_cast_count = 0
187        for graph_input in graph.input:
188            new_node, removed_nodes = self.change_graph_input_type(graph_input, TensorProto.INT32)
189            if new_node:
190                add_cast_count += 1
191            remove_cast_count += len(removed_nodes)
192        logger.info(
193            f"Graph inputs are changed to int32. Added {add_cast_count} Cast nodes, and removed {remove_cast_count} Cast nodes."
194        )
195
196    def use_dynamic_axes(self, dynamic_batch_dim="batch_size", dynamic_seq_len="max_seq_len"):
197        """
198        Update input and output shape to use dynamic axes.
199        """
200        bert_graph_inputs = self.get_graph_inputs_from_fused_nodes(
201            casted=True
202        ) + self.get_graph_inputs_from_fused_nodes(casted=False)
203
204        for input in self.model.graph.input:
205            if input.name in bert_graph_inputs:
206                dim_proto = input.type.tensor_type.shape.dim[0]
207                dim_proto.dim_param = dynamic_batch_dim
208                if dynamic_seq_len is not None:
209                    dim_proto = input.type.tensor_type.shape.dim[1]
210                    dim_proto.dim_param = dynamic_seq_len
211
212        for output in self.model.graph.output:
213            dim_proto = output.type.tensor_type.shape.dim[0]
214            dim_proto.dim_param = dynamic_batch_dim
215
216    def preprocess(self):
217        self.adjust_reshape_and_expand()
218        return
219
220    def adjust_reshape_and_expand(self):
221        nodes_to_remove = []
222        for node in self.nodes():
223            if node.op_type == "Reshape":
224                # Clean up unnecessary reshape nodes.
225                # Find reshape nodes with no actually data in "shape" attribute and remove.
226                reshape_shape = self.get_constant_value(node.input[1])
227                if reshape_shape is not None and reshape_shape.size == 0:
228                    nodes_to_remove.extend([node])
229                    self.replace_input_of_all_nodes(node.output[0], node.input[0])
230                    continue
231
232                # Find path "Slice" -> "Reshape" -> "Expand" -> "Expand" -> current "Reshape", simplify the graph by
233                # changing current reshape's input to output of slice.
234                reshape_path = self.match_parent_path(
235                    node,
236                    ["Expand", "Expand", "Reshape", "Slice"],
237                    [0, 0, 0, 0],
238                    self.output_name_to_node(),
239                )
240                if reshape_path is not None:
241                    expand_node = reshape_path[-3]
242                    expand_shape_value = self.get_constant_value(expand_node.input[1])
243
244                    reshape_before_expand = reshape_path[-2]
245                    shape_value = self.get_constant_value(reshape_before_expand.input[1])
246
247                    slice_node = reshape_path[-1]
248                    if (
249                        expand_shape_value is not None
250                        and shape_value is not None
251                        and len(expand_shape_value) == 2
252                        and len(shape_value) == 1
253                        and expand_shape_value[1] == shape_value[0]
254                    ):
255                        node.input[0] = slice_node.output[0]
256
257        if nodes_to_remove:
258            self.remove_nodes(nodes_to_remove)
259            logger.info(f"Removed Reshape and Expand count: {len(nodes_to_remove)}")
260
261    def clean_graph(self):
262        output_name_to_node = self.output_name_to_node()
263        nodes_to_remove = []
264        for node in self.nodes():
265            # Before:
266            #  input_ids --> Shape --> Gather(indices=0) --> Unsqueeze ------+
267            #          |                                                     |
268            #          |                                                     v
269            #          +----> Shape --> Gather(indices=1) --> Unsqueeze--->  Concat --> ConstantOfShape -->Cast --> EmbedLayerNormaliation/ReduceSum
270            # After (Concat path simplified, Cast merged into ConstantOfShape):
271            #  input_ids --> Shape --> ConstantOfShape --> EmbedLayerNormalization/ReduceSum
272            op_input_id = {"EmbedLayerNormalization": 1, "ReduceSum": 0, "Attention": 3}
273            if node.op_type in op_input_id:
274                i = op_input_id[node.op_type]
275                parent_nodes = self.match_parent_path(
276                    node,
277                    [
278                        "Cast",
279                        "ConstantOfShape",
280                        "Concat",
281                        "Unsqueeze",
282                        "Gather",
283                        "Shape",
284                    ],
285                    [i, 0, 0, 0, 0, 0],
286                    output_name_to_node,
287                )
288                if parent_nodes is not None:
289                    (
290                        cast,
291                        constant_of_shape,
292                        concat,
293                        unsqueeze,
294                        gather,
295                        shape,
296                    ) = parent_nodes
297                    if shape.input[0] == self.graph().input[0].name:
298                        constant_of_shape.input[0] = shape.output[0]
299
300                        # Merge ConstantOfShape → Cast: update the value attribute dtype
301                        # so ConstantOfShape directly produces the target type.
302                        cast_to_type = OnnxModel.get_node_attribute(cast, "to")
303                        cos_tensor = OnnxModel.get_node_attribute(constant_of_shape, "value")
304                        if cast_to_type is not None and cos_tensor is not None:
305                            fill_val = numpy_helper.to_array(cos_tensor).flat[0]
306                            np_dtype = helper.tensor_dtype_to_np_dtype(cast_to_type)
307                            new_val = numpy_helper.from_array(np.array([fill_val], dtype=np_dtype))
308                            for i, attr in enumerate(constant_of_shape.attribute):
309                                if attr.name == "value":
310                                    constant_of_shape.attribute[i].CopyFrom(helper.make_attribute("value", new_val))
311                                    break
312                            self.replace_input_of_all_nodes(cast.output[0], constant_of_shape.output[0])
313                            nodes_to_remove.append(cast)
314
315                        output_name_to_node = self.output_name_to_node()
316
317            if node.op_type == "Attention":
318                # Before (Cast present or already merged into ConstantOfShape):
319                #   input_ids --> Shape --> ConstantOfShape [--> Cast] --> ReduceSum --> Attention
320                # After:
321                #   remove this path, and remove the optional mask_index input of Attention node.
322                parent_nodes = self.match_parent_path(
323                    node,
324                    ["ReduceSum", "Cast", "ConstantOfShape", "Shape"],
325                    [3, 0, 0, 0],
326                    output_name_to_node,
327                )
328                if parent_nodes is None:
329                    # Also try merged pattern (Cast already folded into ConstantOfShape).
330                    parent_nodes = self.match_parent_path(
331                        node,
332                        ["ReduceSum", "ConstantOfShape", "Shape"],
333                        [3, 0, 0],
334                        output_name_to_node,
335                    )
336                if parent_nodes is not None:
337                    if parent_nodes[-1].input[0] == self.graph().input[0].name:
338                        attention_node = helper.make_node(
339                            "Attention",
340                            inputs=node.input[0 : len(node.input) - 1],
341                            outputs=node.output,
342                            name=node.name + "_remove_mask",
343                        )
344                        attention_node.domain = "com.microsoft"
345                        attention_node.attribute.extend([helper.make_attribute("num_heads", self.num_heads)])
346                        self.add_node(attention_node, self.get_graph_by_node(node).name)
347                        nodes_to_remove.append(node)
348        self.remove_nodes(nodes_to_remove)
349
350    def postprocess(self):
351        self.clean_graph()
352        self.prune_graph()
353
354    def optimize(self, options: FusionOptions | None = None, add_dynamic_axes: bool = False):
355        if (options is not None) and not options.enable_shape_inference:
356            self.disable_shape_inference()
357
358        self.utils.remove_identity_nodes()
359
360        # Remove cast nodes that having same data type of input and output based on symbolic shape inference.
361        self.utils.remove_useless_cast_nodes()
362
363        # Apply any missed constant-folding model optimizations (e.g. for Dynamo-exported models)
364        self.fuse_constant_fold()
365
366        if (options is None) or options.enable_layer_norm:
367            self.fuse_layer_norm()
368            self.fuse_simplified_layer_norm()
369
370        if (options is None) or options.enable_gelu:
371            self.fuse_gelu()
372
373        self.preprocess()
374
375        self.fuse_reshape()
376
377        if (options is None) or options.enable_skip_layer_norm:
378            self.fuse_skip_layer_norm(options.enable_shape_inference)
379            self.fuse_skip_simplified_layer_norm()
380
381        if (options is None) or options.enable_rotary_embeddings:
382            self.fuse_rotary_embeddings()
383
384        if options is not None:
385            self.attention_mask.set_mask_format(options.attention_mask_format)
386            if options.use_multi_head_attention and not isinstance(self.attention_fusion, FusionBartAttention):
387                self.attention_fusion = FusionAttention(
388                    self,
389                    self.hidden_size,
390                    self.num_heads,
391                    self.attention_mask,
392                    options.use_multi_head_attention,
393                )
394
395        if (options is None) or options.enable_attention:
396            self.fuse_attention()
397
398        # Perform the MatMul fusion after the Attention fusion as we do not
399        # want to fuse the MatMuls inside the Attention subgraphs
400        if (options is None) or options.enable_qordered_matmul:
401            self.fuse_qordered_mamtul()
402
403        self.fuse_shape()
404
405        if (options is None) or options.enable_embed_layer_norm:
406            use_mask_index = options.attention_mask_format == AttentionMaskFormat.MaskIndexEnd
407            self.fuse_embed_layer(use_mask_index)
408
409        # Remove reshape nodes that having same shape of input and output based on symbolic shape inference.
410        self.utils.remove_useless_reshape_nodes()
411
412        self.postprocess()
413
414        # Bias fusion is done after postprocess to avoid extra Reshape between bias and Gelu/FastGelu/SkipLayerNormalization
415        if (options is None) or options.enable_bias_gelu:
416            # Fuse Gelu and Add Bias before it.
417            self.fuse_bias_gelu(is_fastgelu=True)
418            self.fuse_bias_gelu(is_fastgelu=False)
419
420        if (options is None) or options.enable_bias_skip_layer_norm:
421            # Fuse SkipLayerNormalization and Add Bias before it.
422            self.fuse_add_bias_skip_layer_norm()
423
424        if options is not None and options.enable_gelu_approximation:
425            self.gelu_approximation()
426
427        if options is not None and options.enable_gemm_fast_gelu:
428            self.fuse_gemm_fast_gelu()
429
430        self.remove_unused_constant()
431
432        # Use symbolic batch dimension in input and output.
433        if add_dynamic_axes:
434            self.use_dynamic_axes()
435
436        logger.info(f"opset version: {self.get_opset_version()}")
437
438    def get_fused_operator_statistics(self):
439        """
440        Returns node count of fused operators.
441        """
442        op_count = {}
443        ops = [
444            "EmbedLayerNormalization",
445            "Attention",
446            "MultiHeadAttention",
447            "Gelu",
448            "FastGelu",
449            "BiasGelu",
450            "GemmFastGelu",
451            "LayerNormalization",
452            "SimplifiedLayerNormalization",
453            "SkipLayerNormalization",
454            "SkipSimplifiedLayerNormalization",
455            "RotaryEmbedding",
456        ]
457        q_ops = [
458            "QOrderedAttention",
459            "QOrderedGelu",
460            "QOrderedLayerNormalization",
461            "QOrderedMatMul",
462        ]
463        for op in ops + q_ops:
464            nodes = self.get_nodes_by_op_type(op)
465            op_count[op] = len(nodes)
466
467        logger.info(f"Optimized operators: {op_count}")
468        return op_count
469
470    def is_fully_optimized(self, fused_op_count=None):
471        """
472        Returns True when the model is fully optimized.
473        """
474        if fused_op_count is None:
475            fused_op_count = self.get_fused_operator_statistics()
476
477        def op_count(op_name: str):
478            return fused_op_count.get(op_name) or 0
479
480        embed = op_count("EmbedLayerNormalization")
481        attention = op_count("Attention") + op_count("MultiHeadAttention") + op_count("QOrderedAttention")
482        gelu = op_count("Gelu") + op_count("BiasGelu") + op_count("FastGelu")
483        layer_norm = op_count("LayerNormalization") + op_count("SkipLayerNormalization")
484        simple_layer_norm = op_count("SimplifiedLayerNormalization") + op_count("SkipSimplifiedLayerNormalization")
485
486        is_perfect = (
487            (embed > 0)
488            and (attention > 0)
489            and (attention == gelu)
490            and ((layer_norm >= 2 * attention) or (simple_layer_norm >= 2 * attention))
491        )
492
493        if layer_norm == 0:
494            logger.debug("Layer Normalization not fused")
495
496        if simple_layer_norm == 0:
497            logger.debug("Simple Layer Normalization not fused")
498
499        if gelu == 0:
500            logger.debug("Gelu (or FastGelu) not fused")
501
502        if embed == 0:
503            logger.debug("EmbedLayerNormalization not fused")
504
505        if attention == 0:
506            logger.warning("Attention (or MultiHeadAttention) not fused")
507
508        return is_perfect
509
510    def convert_to_packing_mode(self, use_symbolic_shape_infer: bool = False):
511        packing_mode = PackingMode(self)
512        packing_mode.convert(use_symbolic_shape_infer)
513 
codekingpro/portable-devtools · Team Ai