Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_model_bert_tf.py589 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6import logging
7
8import numpy as np
9import onnx
10from onnx import TensorProto, helper, numpy_helper
11from onnx_model_bert import BertOnnxModel
12
13logger = logging.getLogger(__name__)
14
15
16class BertOnnxModelTF(BertOnnxModel):
17    def __init__(self, model, num_heads, hidden_size):
18        super().__init__(model, num_heads, hidden_size)
19
20    def remove_identity(self):
21        nodes_to_remove = []
22        for node in self.nodes():
23            if node.op_type == "Identity":
24                if not self.find_graph_output(node.output[0]):
25                    self.replace_input_of_all_nodes(node.output[0], node.input[0])
26                    nodes_to_remove.append(node)
27        self.remove_nodes(nodes_to_remove)
28        logger.info(f"Removed Identity count: {len(nodes_to_remove)}")
29
30    def match_mask_path(self, add_or_sub_before_softmax):
31        mask_nodes = self.match_parent_path(
32            add_or_sub_before_softmax,
33            ["Mul", "Sub", "Reshape", "Cast"],
34            [1, None, 1, 0],
35        )
36        if mask_nodes is not None:
37            return mask_nodes
38
39        mask_nodes = self.match_parent_path(
40            add_or_sub_before_softmax,
41            ["Mul", "Sub", "Cast", "Slice", "Unsqueeze"],
42            [1, 0, 1, 0, 0],
43        )
44        if mask_nodes is not None:
45            return mask_nodes
46
47        mask_nodes = self.match_parent_path(
48            add_or_sub_before_softmax,
49            ["Mul", "Sub", "Cast", "Unsqueeze", "Unsqueeze"],
50            [1, None, 1, 0, 0],
51        )
52
53        return mask_nodes
54
55    def get_2d_initializers_from_parent_subgraphs(self, current_node):
56        """
57        Find initializers that is 2D. Returns a dictionary with name as key and shape as value.
58        """
59        parent_nodes = self.get_parent_subgraph_nodes(current_node, [])
60        initializers = {}
61        for node in parent_nodes:
62            for input in node.input:
63                initializer = self.get_initializer(input)
64                if initializer:
65                    temp = numpy_helper.to_array(initializer)
66                    if len(temp.shape) == 2:
67                        initializers[initializer.name] = temp.shape
68
69        return initializers
70
71    def find_segment_ids(self, segment_embedding, input_ids):
72        input_name_to_nodes = self.input_name_to_nodes()
73        if segment_embedding not in input_name_to_nodes:
74            return None
75
76        nodes = input_name_to_nodes[segment_embedding]
77        if len(nodes) != 1:
78            return None
79
80        graph_inputs = self.get_graph_inputs(nodes[0], recursive=True)
81        if len(graph_inputs) > 1:
82            print("Found multiple candidates of segment_ids", graph_inputs)
83            return None
84        # Find segment ids in graph inputs. The segment id input must not be the same as input_ids.
85        if len(graph_inputs) == 1 and graph_inputs[0] != input_ids:
86            return graph_inputs[0]
87
88        # If the segment id candidate is the same as the input_ids, try to assign alternative segment ids and simplify the graph if needed.
89        segment_ids = nodes[0].input[1]
90        _, segment_id_path, _ = self.match_parent_paths(
91            nodes[0],
92            [
93                (
94                    ["ConstantOfShape", "Cast", "Concat", "Slice", "Cast", "Shape"],
95                    [1, 0, 0, 0, 0, 0],
96                ),
97                (
98                    [
99                        "ConstantOfShape",
100                        "Cast",
101                        "Concat",
102                        "Unsqueeze",
103                        "Squeeze",
104                        "Slice",
105                        "Cast",
106                        "Shape",
107                    ],
108                    [1, 0, 0, 0, 0, 0, 0, 0],
109                ),
110            ],
111            None,
112        )
113
114        if segment_id_path and input_ids and input_ids == segment_id_path[-1].input[0]:
115            logger.debug("Simplify semgent id path...")
116            constantofshape_node = segment_id_path[0]
117            graph_name = self.get_graph_by_node(constantofshape_node).name
118            self.add_node(
119                helper.make_node("Shape", inputs=[input_ids], outputs=["input_shape"]),
120                graph_name,
121            )
122            constantofshape_value = helper.get_attribute_value(constantofshape_node.attribute[0])
123            self.add_node(
124                helper.make_node(
125                    "ConstantOfShape",
126                    inputs=["input_shape"],
127                    outputs=["zeros_for_input_shape"],
128                    value=constantofshape_value,
129                ),
130                graph_name,
131            )
132            segment_ids = "zeros_for_input_shape"
133        return segment_ids
134
135    def find_input_ids(self, word_embedding):
136        input_name_to_nodes = self.input_name_to_nodes()
137        if word_embedding not in input_name_to_nodes:
138            return None
139
140        nodes = input_name_to_nodes[word_embedding]
141        if len(nodes) != 1:
142            return None
143
144        graph_inputs = self.get_graph_inputs(nodes[0], recursive=True)
145        if len(graph_inputs) == 1:
146            return graph_inputs[0]
147
148        print("Found multiple candidates of input_ids", graph_inputs)
149        return None
150
151    def find_mask_input(self, excluded_graph_inputs):
152        for node in self.nodes():
153            if node.op_type == "Softmax":
154                mask_path = self.match_parent_path(
155                    node,
156                    ["Add", "Mul", "Sub", "Cast", "Slice", "Unsqueeze"],
157                    [0, 1, None, 1, 0, 0],
158                )
159                if mask_path is None:
160                    continue
161                (
162                    add_node,
163                    mul_node,
164                    sub_node,
165                    cast_node,
166                    slice_node,
167                    unsqueeze_node,
168                ) = mask_path
169                if self.has_constant_input(mul_node, -10000) and self.has_constant_input(sub_node, 1):
170                    graph_inputs = self.get_graph_inputs(sub_node, recursive=True)
171                    inputs = [input for input in graph_inputs if input not in excluded_graph_inputs]
172                    if len(inputs) > 1:
173                        print("Found multiple candidates of mask input", inputs)
174                        return None
175                    if len(inputs) == 1:
176                        return inputs[0]
177                    # Duplicated input found. Try to simplify the graph.
178                    path_to_be_simplified = self.match_parent_path(
179                        mask_path[-1],
180                        [
181                            "ConstantOfShape",
182                            "Cast",
183                            "Concat",
184                            "Unsqueeze",
185                            "Squeeze",
186                            "Slice",
187                            "Cast",
188                            "Shape",
189                        ],
190                        [0, 0, 0, 0, 0, 0, 0, 0],
191                    )
192                    duplicated_inputs = [input for input in graph_inputs if input in excluded_graph_inputs]
193                    # Simplify graph for dynamic axes.
194                    if (
195                        path_to_be_simplified
196                        and duplicated_inputs
197                        and len(duplicated_inputs) == 1
198                        and duplicated_inputs[0] == path_to_be_simplified[-1].input[0]
199                    ):
200                        logger.debug("Simplify semgent id path...")
201                        constantofshape_node = path_to_be_simplified[0]
202                        constantofshape_value = helper.get_attribute_value(constantofshape_node.attribute[0])
203                        graph_name = self.get_graph_by_node(constantofshape_node).name
204                        self.add_node(
205                            helper.make_node(
206                                "Shape",
207                                inputs=[duplicated_inputs[0]],
208                                outputs=["input_shape_for_mask"],
209                            ),
210                            graph_name,
211                        )
212                        self.add_node(
213                            helper.make_node(
214                                "ConstantOfShape",
215                                inputs=["input_shape_for_mask"],
216                                outputs=[unsqueeze_node.input[0]],
217                                value=constantofshape_value,
218                            ),
219                            graph_name,
220                        )
221                    return unsqueeze_node.input[0]
222        return None
223
224    def create_embedding_subgraph(self, normalize_node, word_embedding, segment_embedding, position_embedding):
225        input_ids = self.find_input_ids(word_embedding)
226        if input_ids is None:
227            logger.info("Failed to find input_ids. Cannot fuse embedding layer.")
228            return False
229
230        segment_ids = self.find_segment_ids(segment_embedding, input_ids)
231        if segment_ids is None:
232            logger.info("Failed to find segment_ids. Cannot fuse embedding layer.")
233            return False
234
235        mask_input = self.find_mask_input([segment_ids, input_ids])
236        if mask_input is None:
237            logger.info("Failed to find input_mask. Cannot fuse embedding layer.")
238            return False
239
240        self.bert_inputs = [input_ids, segment_ids, mask_input]
241
242        mask_index = self.create_node_name("mask_index")
243        self.attention_mask.set_mask_indice(mask_input, mask_index)
244
245        if self.find_graph_input(input_ids).type.tensor_type.elem_type != TensorProto.INT32:
246            casted, input_ids = self.utils.cast_graph_input_to_int32(input_ids)
247
248        if self.find_graph_input(segment_ids):
249            casted, segment_ids = self.utils.cast_graph_input_to_int32(segment_ids)
250        else:
251            segment_ids, segment_id_cast_node = self.utils.cast_input_to_int32(segment_ids)
252
253        if self.find_graph_input(mask_input):
254            casted, mask_input = self.utils.cast_graph_input_to_int32(mask_input)
255        else:
256            mask_input, mask_input_cast_node = self.utils.cast_input_to_int32(mask_input)
257
258        embed_output = self.create_node_name("embed_output")
259        embed_node = onnx.helper.make_node(
260            "EmbedLayerNormalization",
261            inputs=[
262                input_ids,
263                segment_ids,
264                word_embedding,
265                position_embedding,
266                segment_embedding,
267                normalize_node.input[1],  # gamma
268                normalize_node.input[2],  # beta
269                mask_input,
270            ],
271            outputs=[embed_output, mask_index],
272            name="EmbedLayer",
273        )
274        embed_node.domain = "com.microsoft"
275        self.replace_input_of_all_nodes(normalize_node.output[0], embed_output)
276        self.add_node(embed_node, self.get_graph_by_node(normalize_node).name)
277
278    def process_embedding(self):
279        """
280        Automatically detect word, segment and position embeddings.
281        """
282        logger.info("start processing embedding layer...")
283        output_name_to_node = self.output_name_to_node()
284
285        layer_norm_nodes = self.get_nodes_by_op_type("LayerNormalization")
286        for layer_norm_node in layer_norm_nodes:
287            pos_embed_path = self.match_parent_path(
288                layer_norm_node,
289                ["Add", "Reshape", "Slice"],
290                [0, 1, 0],
291                output_name_to_node,
292            )
293            if pos_embed_path is None:
294                continue
295
296            add_node, reshape_node, slice_node = pos_embed_path
297            initializer = self.get_initializer(slice_node.input[0])
298            if initializer is None:
299                continue
300
301            temp = numpy_helper.to_array(initializer)
302            if len(temp.shape) == 2:
303                logger.info(f"Found position embedding. name:{initializer.name}, shape:{temp.shape}")
304                position_embedding = initializer.name
305            else:
306                logger.info(f"Failed to find position embedding. name:{initializer.name}, shape:{temp.shape}")
307                return
308
309            first_parent = self.get_parent(add_node, 0, output_name_to_node)
310            if first_parent is not None and first_parent.op_type == "Add":
311                embeddings = self.get_2d_initializers_from_parent_subgraphs(first_parent)
312                if len(embeddings) != 2:
313                    logger.warning(
314                        f"Failed to find two embeddings (word and segment) from Add node. Found {embeddings}"
315                    )
316                    return
317
318                word_embedding = None
319                segment_embedding = None
320                for name, shape in embeddings.items():
321                    if shape[0] == 2:
322                        segment_embedding = name
323                        logger.info(f"Found segment embedding. name:{name}, shape:{shape}")
324                    else:
325                        word_embedding = name
326                        logger.info(f"Found words embedding. name:{name}, shape:{shape}")
327
328                if word_embedding is None or segment_embedding is None:
329                    logger.info("Failed to find both word and segment embedding")
330                    return
331
332                logger.info("Create Embedding node")
333                self.create_embedding_subgraph(
334                    layer_norm_node,
335                    word_embedding,
336                    segment_embedding,
337                    position_embedding,
338                )
339                # Prune graph to remove those original embedding nodes.
340                self.prune_graph()
341                break
342
343    def check_attention_input(self, matmul_q, matmul_k, matmul_v, parent, output_name_to_node):
344        for x in [matmul_q, matmul_k, matmul_v]:
345            root_input = x.input[0]
346            root_node = output_name_to_node[root_input]
347            if root_node == parent:
348                continue
349            logger.debug(f"Check attention input failed:{root_input}, {parent.output[0]}")
350            return False
351
352        return True
353
354    def fuse_attention(self):
355        output_name_to_node = self.output_name_to_node()
356
357        nodes_to_remove = []
358        attention_count = 0
359
360        start_nodes = []
361        skip_layer_norm_nodes = self.get_nodes_by_op_type("SkipLayerNormalization")
362        layer_norm_nodes = self.get_nodes_by_op_type("LayerNormalization")
363        # Sometimes we can not fuse skiplayernormalization since the add before layernorm has an output that used by nodes outside skiplayernorm
364        # Conceptually we treat add before layernorm as skiplayernorm node since they share the same pattern
365        start_nodes.extend(skip_layer_norm_nodes)
366        start_nodes.extend(layer_norm_nodes)
367
368        for normalize_node in start_nodes:
369            graph_name = self.get_graph_by_node(normalize_node).name
370            # SkipLayerNormalization has two inputs, and one of them is the root input for attention.
371            if normalize_node.op_type == "LayerNormalization":
372                add_before_layernorm = self.match_parent(normalize_node, "Add", 0)
373                if add_before_layernorm is not None:
374                    normalize_node = add_before_layernorm  # noqa: PLW2901
375                else:
376                    continue
377            parent = self.get_parent(normalize_node, 1)
378            if parent is None or parent.op_type not in [
379                "SkipLayerNormalization",
380                "LayerNormalization",
381                "Reshape",
382            ]:
383                parent = self.get_parent(normalize_node, 0)
384                if parent is None or parent.op_type not in [
385                    "SkipLayerNormalization",
386                    "LayerNormalization",
387                    "Reshape",
388                ]:
389                    logger.debug("Failed to match parent of normalize_node")
390                    continue
391
392            qkv_nodes = self.match_parent_path(
393                normalize_node,
394                ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
395                [0, 0, 0, 0, 0],
396            )
397            if qkv_nodes is None:
398                qkv_nodes = self.match_parent_path(
399                    normalize_node,
400                    ["MatMul", "Reshape", "Transpose", "MatMul"],
401                    [1, 0, 0, 0],
402                )
403                if qkv_nodes is None:
404                    qkv_nodes = self.match_parent_path(normalize_node, ["Add", "Einsum", "Einsum"], [0, 0, 0])
405                    if qkv_nodes is None:
406                        logger.debug("Failed to match qkv nodes")
407                        continue
408
409            matmul_qkv = qkv_nodes[-1]
410            v_nodes = self.match_parent_path(matmul_qkv, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, 0])
411            if v_nodes is None:
412                v_nodes = self.match_parent_path(matmul_qkv, ["Add", "Einsum"], [1, 0])
413                if v_nodes is None:
414                    logger.debug("Failed to match v path")
415                    continue
416
417            add_v = v_nodes[-2]
418            matmul_v = v_nodes[-1]
419            qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Add", "Mul", "MatMul"], [0, 0, 0, 0])
420            if qk_nodes is None:
421                qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Add", "Einsum"], [0, 0, 0])
422                if qk_nodes is None:
423                    logger.debug("Failed to match qk_paths")
424                    continue
425            matmul_qk = qk_nodes[-1]
426
427            q_nodes = self.match_parent_path(matmul_qk, ["Transpose", "Reshape", "Add", "MatMul"], [0, 0, 0, 0])
428            if q_nodes is None:
429                q_nodes = self.match_parent_path(matmul_qk, ["Add", "Einsum"], [0, 0])
430                if q_nodes is None:
431                    logger.debug("Failed to match q path")
432                    continue
433
434            add_q = q_nodes[-2]
435            matmul_q = q_nodes[-1]
436
437            k_nodes = self.match_parent_path(matmul_qk, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, 0])
438            if k_nodes is None:
439                k_nodes = self.match_parent_path(matmul_qk, ["Mul", "Add", "Einsum"], [1, 0, 0])
440                if k_nodes is None:
441                    logger.debug("Failed to match k path")
442                    continue
443            add_k = k_nodes[-2]
444            matmul_k = k_nodes[-1]
445
446            mask_nodes = self.match_mask_path(qk_nodes[1])
447
448            if mask_nodes is None:
449                logger.debug("Cannot find mask_nodes.")
450                continue
451
452            if not self.has_constant_input(mask_nodes[1], 1):
453                logger.debug("Sub node expected to have an input with constant value 1.0.")
454                continue
455
456            # add a squeeze node to convert a 3-d mask to 2-d
457            squeeze_node = self.match_parent_path(mask_nodes[-1], ["Squeeze"], [0]) or self.match_parent_path(
458                mask_nodes[-1], ["Expand"], [0]
459            )
460            squeeze_node_name = "Squeeze_3d_to_2d_mask"
461            squeeze_output_name = squeeze_node_name + "_output"
462            if squeeze_node is None and len(mask_nodes) == 5 and self.find_graph_input(mask_nodes[-1].input[0]) is None:
463                mask_input = mask_nodes[-1].input[1]
464                self.add_node(
465                    helper.make_node(
466                        "Squeeze",
467                        [mask_input],
468                        [squeeze_output_name],
469                        squeeze_node_name,
470                        axes=[1],
471                    ),
472                    graph_name,
473                )
474                mask_nodes[-1].input[0] = squeeze_output_name
475
476            is_same_root = self.check_attention_input(matmul_q, matmul_k, matmul_v, parent, output_name_to_node)
477            if is_same_root:
478                mask_index = self.attention_mask.process_mask(mask_nodes[-1].input[0])
479                logger.debug("Create an Attention node.")
480
481                # For tf models, q and v are flipped.
482                attention_node = self.attention_fusion.create_attention_node(
483                    mask_index=mask_index,
484                    q_matmul=matmul_k,
485                    k_matmul=matmul_q,
486                    v_matmul=matmul_v,
487                    q_add=add_k,
488                    k_add=add_q,
489                    v_add=add_v,
490                    num_heads=self.num_heads,
491                    hidden_size=self.hidden_size,
492                    first_input=parent.output[0],
493                    output=qkv_nodes[2].output[0],
494                )
495                if attention_node is None:
496                    continue
497
498                if qkv_nodes[1].op_type == "Einsum":
499                    # add reshape before einsum
500                    tensor = helper.make_tensor(
501                        name=qkv_nodes[1].name + "_newshape",
502                        data_type=TensorProto.INT64,
503                        dims=[4],
504                        vals=np.int64(
505                            [
506                                [
507                                    0,
508                                    0,
509                                    self.num_heads,
510                                    int(self.hidden_size / self.num_heads),
511                                ]
512                            ]
513                        ).tobytes(),
514                        raw=True,
515                    )
516                    self.add_initializer(tensor, graph_name)
517                    reshape_ = helper.make_node(
518                        "Reshape",
519                        inputs=[
520                            attention_node.output[0],
521                            qkv_nodes[1].name + "_newshape",
522                        ],
523                        outputs=[qkv_nodes[1].name + "_reshape_output"],
524                        name=qkv_nodes[1].name + "_reshape",
525                    )
526                    qkv_nodes[1].input[0] = qkv_nodes[1].name + "_reshape_output"
527                    self.add_node(reshape_, graph_name)
528                if parent.op_type == "Reshape":
529                    # Temporary work around: we require the skiplayernorm and attention op be fed with 3-d input
530                    hidden_size = numpy_helper.to_array(self.get_initializer(parent.input[1]))[1]
531                    tensor = helper.make_tensor(
532                        name=parent.name + "_modified",
533                        data_type=TensorProto.INT64,
534                        dims=[3],
535                        vals=np.int64([[1, -1, hidden_size]]).tobytes(),
536                        raw=True,
537                    )
538                    self.add_initializer(tensor, graph_name)
539                    parent.input[1] = parent.name + "_modified"
540
541                self.add_node(attention_node, graph_name)
542                attention_count += 1
543
544                nodes_to_remove.extend(qkv_nodes[2:])
545                nodes_to_remove.extend(qk_nodes)
546                nodes_to_remove.extend(q_nodes)
547                nodes_to_remove.extend(k_nodes)
548                nodes_to_remove.extend(v_nodes)
549                nodes_to_remove.extend(mask_nodes)
550            else:
551                logger.debug("Root node not matched.")
552                continue
553        self.remove_nodes(nodes_to_remove)
554        self.update_graph()
555        logger.info(f"Fused Attention count:{attention_count}")
556
557    def preprocess(self):
558        self.remove_identity()
559        self.process_embedding()
560        self.skip_reshape()
561
562    def skip_reshape(self):
563        count = 0
564        reshape_nodes = self.get_nodes_by_op_type("Reshape")
565        for reshape_node in reshape_nodes:
566            parent = self.get_parent(reshape_node, 0)
567            if parent is not None and parent.op_type == "Reshape":
568                reshape_node.input[0] = parent.input[0]
569                count += 1
570
571        if count > 0:
572            logger.info(f"Skip consequent Reshape count: {count}")
573
574    def remove_reshape_before_first_attention(self):
575        attention_nodes = self.get_nodes_by_op_type("Attention")
576        for attention_node in attention_nodes:
577            path = self.match_parent_path(attention_node, ["Reshape", "EmbedLayerNormalization"], [0, 0])
578            if path is None:
579                continue
580            logger.info("Remove Reshape before first Attention node.")
581            reshape, _ = path
582            self.replace_input_of_all_nodes(reshape.output[0], reshape.input[0])
583            self.remove_node(reshape)
584            break
585
586    def postprocess(self):
587        self.remove_reshape_before_first_attention()
588        self.prune_graph()
589 
codekingpro/portable-devtools · Team Ai