Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_model_bert_keras.py475 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 onnx
9from onnx import numpy_helper
10from onnx_model_bert_tf import BertOnnxModelTF
11
12logger = logging.getLogger(__name__)
13
14
15class BertOnnxModelKeras(BertOnnxModelTF):
16    def __init__(self, model, num_heads, hidden_size):
17        super().__init__(model, num_heads, hidden_size)
18
19    def match_mask_path(self, add_or_sub_before_softmax):
20        mask_nodes = self.match_parent_path(
21            add_or_sub_before_softmax,
22            ["Mul", "Sub", "Reshape", "Cast"],
23            [1, None, 1, 0],
24        )
25        if mask_nodes is not None:
26            return mask_nodes
27
28        mask_nodes = self.match_parent_path(
29            add_or_sub_before_softmax,
30            ["Mul", "Sub", "Cast", "Slice", "Unsqueeze"],
31            [1, 1, 1, 0, 0],
32        )
33        if mask_nodes is not None:
34            return mask_nodes
35
36        mask_nodes = self.match_parent_path(
37            add_or_sub_before_softmax,
38            ["Mul", "Sub", "Cast", "Unsqueeze", "Unsqueeze"],
39            [1, None, 1, 0, 0],
40        )
41        return mask_nodes
42
43    def check_attention_input(self, matmul_q, matmul_k, matmul_v, parent, output_name_to_node):
44        reshape_nodes = []
45
46        for x in [matmul_q, matmul_k, matmul_v]:
47            root_input = x.input[0]
48            root_node = output_name_to_node[root_input]
49            if root_node == parent:
50                continue
51            if root_node.op_type == "Reshape" and root_node.input[0] == parent.output[0]:
52                reshape_nodes.append(root_node)
53                continue
54            logger.debug(f"Check attention input failed:{root_input}, {parent.output[0]}")
55            return False, []
56
57        return True, reshape_nodes
58
59    def fuse_attention(self):
60        self.input_name_to_nodes()
61        output_name_to_node = self.output_name_to_node()
62
63        nodes_to_remove = []
64        attention_count = 0
65
66        skip_layer_norm_nodes = self.get_nodes_by_op_type("SkipLayerNormalization")
67        for normalize_node in skip_layer_norm_nodes:
68            # SkipLayerNormalization has two inputs, and one of them is the root input for attention.
69            parent = self.get_parent(normalize_node, 0)
70            if parent is None or parent.op_type not in [
71                "SkipLayerNormalization",
72                "EmbedLayerNormalization",
73            ]:
74                if parent.op_type == "Add":
75                    parent = self.get_parent(normalize_node, 1)
76                    if parent is None or parent.op_type not in [
77                        "SkipLayerNormalization",
78                        "EmbedLayerNormalization",
79                    ]:
80                        logger.debug(f"First input for skiplayernorm: {parent.op_type if parent is not None else None}")
81                        continue
82                else:
83                    logger.debug(f"First input for skiplayernorm: {parent.op_type if parent is not None else None}")
84                    continue
85            else:
86                # TODO: shall we add back the checking of children op types.
87                pass
88
89            qkv_nodes = self.match_parent_path(
90                normalize_node,
91                ["Add", "Reshape", "MatMul", "Reshape", "Transpose", "MatMul"],
92                [None, 0, 0, 0, 0, 0],
93            )
94            if qkv_nodes is None:
95                logger.debug("Failed to match qkv nodes")
96                continue
97            (
98                add,
99                extra_reshape_0,
100                matmul,
101                reshape_qkv,
102                transpose_qkv,
103                matmul_qkv,
104            ) = qkv_nodes
105            logger.debug("Matched qkv nodes")
106
107            v_nodes = self.match_parent_path(
108                matmul_qkv,
109                ["Transpose", "Reshape", "Add", "Reshape", "MatMul"],
110                [1, 0, 0, 0, 0],
111            )
112            if v_nodes is None:
113                logger.debug("Failed to match v path")
114                continue
115            (transpose_v, reshape_v, add_v, extra_reshape_1, matmul_v) = v_nodes
116
117            qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Sub", "MatMul"], [0, 0, 0])
118            if qk_nodes is not None:
119                (softmax_qk, sub_qk, matmul_qk) = qk_nodes
120                q_nodes = self.match_parent_path(
121                    matmul_qk,
122                    ["Mul", "Transpose", "Reshape", "Add", "Reshape", "MatMul"],
123                    [0, None, 0, 0, 0, 0],
124                )
125                if q_nodes is not None:
126                    (
127                        mul_q,
128                        transpose_q,
129                        reshape_q,
130                        add_q,
131                        extra_reshape_2,
132                        matmul_q,
133                    ) = q_nodes
134
135            else:
136                qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Add", "Mul", "MatMul"], [0, 0, 0, None])
137                if qk_nodes is None:
138                    qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Add", "Div", "MatMul"], [0, 0, 0, None])
139                    if qk_nodes is None:
140                        logger.debug("Failed to match qk path")
141                        continue
142                (softmax_qk, add_qk, mul_qk, matmul_qk) = qk_nodes
143
144                q_nodes = self.match_parent_path(
145                    matmul_qk,
146                    ["Transpose", "Reshape", "Add", "Reshape", "MatMul"],
147                    [0, 0, 0, 0, 0],
148                )
149                if q_nodes is not None:
150                    (transpose_q, reshape_q, add_q, extra_reshape_2, matmul_q) = q_nodes
151
152            if q_nodes is None:
153                logger.debug("Failed to match q path")
154                continue
155
156            k_nodes = self.match_parent_path(
157                matmul_qk,
158                ["Transpose", "Reshape", "Add", "Reshape", "MatMul"],
159                [1, 0, 0, 0, 0],
160            )
161            if k_nodes is None:
162                logger.debug("Failed to match k path")
163                continue
164            (transpose_k, reshape_k, add_k, extra_reshape_3, matmul_k) = k_nodes
165
166            mask_nodes = self.match_mask_path(qk_nodes[1])
167            if mask_nodes is None:
168                logger.debug("Failed to match mask path")
169                continue
170            if not self.has_constant_input(mask_nodes[1], 1):
171                logger.debug("Sub node expected to have an input with constant value 1.0.")
172                continue
173
174            is_same_root, reshape_nodes = self.check_attention_input(
175                matmul_q, matmul_k, matmul_v, parent, output_name_to_node
176            )
177            if is_same_root:
178                mask_index = self.attention_mask.process_mask(mask_nodes[-1].input[0])
179                logger.debug("Create an Attention node.")
180                attention_node = self.attention_fusion.create_attention_node(
181                    mask_index=mask_index,
182                    q_matmul=matmul_q,
183                    k_matmul=matmul_k,
184                    v_matmul=matmul_v,
185                    q_add=add_q,
186                    k_add=add_k,
187                    v_add=add_v,
188                    num_heads=self.num_heads,
189                    hidden_size=self.hidden_size,
190                    first_input=parent.output[0],
191                    output=reshape_qkv.output[0],
192                )
193                if attention_node is None:
194                    continue
195
196                self.add_node(attention_node)
197                attention_count += 1
198
199                nodes_to_remove.extend([reshape_qkv, transpose_qkv, matmul_qkv])
200                nodes_to_remove.extend(qk_nodes)
201                nodes_to_remove.extend(q_nodes)
202                nodes_to_remove.extend(k_nodes)
203                nodes_to_remove.extend(v_nodes)
204                nodes_to_remove.extend(mask_nodes)
205                nodes_to_remove.extend(reshape_nodes)
206                nodes_to_remove.append(extra_reshape_0)
207                self.replace_node_input(add, extra_reshape_0.output[0], matmul.output[0])
208            else:
209                logger.debug("Root node not matched.")
210                continue
211        self.remove_nodes(nodes_to_remove)
212        self.update_graph()
213        logger.info(f"Fused Attention count:{attention_count}")
214
215    def preprocess(self):
216        self.process_embedding()
217        self.fuse_mask()
218        self.skip_reshape()
219
220    def skip_reshape(self):
221        self.input_name_to_nodes()
222        self.output_name_to_node()
223
224        count = 0
225        reshape_nodes = self.get_nodes_by_op_type("Reshape")
226        for reshape_node in reshape_nodes:
227            parent = self.get_parent(reshape_node, 0)
228            if parent is not None and parent.op_type == "Reshape":
229                reshape_node.input[0] = parent.input[0]
230                count += 1
231
232        if count > 0:
233            logger.info(f"Skip consequent Reshape count: {count}")
234
235    def fuse_embedding(self, node, output_name_to_node):
236        assert node.op_type == "LayerNormalization"
237        logger.debug(f"start fusing embedding from node with output={node.output[0]}...")
238        word_embed_path = self.match_parent_path(node, ["Add", "Add", "Gather"], [0, 0, 0], output_name_to_node)
239        if word_embed_path is None:
240            logger.debug("failed to match word_embed_path")
241            return False
242
243        skip_node, add_node, gather_node = word_embed_path
244
245        word_initializer = self.get_initializer(gather_node.input[0])
246        if word_initializer is None:
247            logger.debug("failed to get word initializer")
248            return False
249
250        temp = numpy_helper.to_array(word_initializer)
251        if len(temp.shape) == 2:
252            logger.info(f"Found word embedding. name:{word_initializer.name}, shape:{temp.shape}")
253            word_embedding = word_initializer.name
254        else:
255            logger.info(f"Failed to find word embedding. name:{word_initializer.name}, shape:{temp.shape}")
256            return False
257
258        pos_initializer = self.get_initializer(add_node.input[1])
259        if pos_initializer is not None:
260            temp = numpy_helper.to_array(pos_initializer)
261            if len(temp.shape) == 3 and temp.shape[0] == 1:
262                tensor = numpy_helper.from_array(temp.reshape((temp.shape[1], temp.shape[2])), "position_embedding")
263                self.add_initializer(tensor)
264                logger.info(f"Found position embedding. name:{pos_initializer.name}, shape:{temp.shape[1:]}")
265                position_embedding = "position_embedding"
266            else:
267                logger.info(f"Failed to find position embedding. name:{pos_initializer.name}, shape:{temp.shape}")
268                return False
269        else:
270            pos_embed_path = self.match_parent_path(add_node, ["Gather", "Slice"], [1, 1], output_name_to_node)
271            if pos_embed_path is None:
272                logger.debug("failed to match pos_embed_path")
273                return False
274
275            pos_gather, pos_slice = pos_embed_path
276            pos_initializer = self.get_initializer(pos_gather.input[0])
277            if pos_initializer is None:
278                logger.debug("failed to get pos initializer")
279                return False
280
281            temp = numpy_helper.to_array(pos_initializer)
282            if len(temp.shape) == 2:
283                logger.info(f"Found word embedding. name:{pos_initializer.name}, shape:{temp.shape}")
284                position_embedding = pos_initializer.name
285            else:
286                logger.info(f"Failed to find position embedding. name:{pos_initializer.name}, shape:{temp.shape}")
287                return False
288
289        gather = self.get_parent(skip_node, 1, output_name_to_node)
290        if gather is None or gather.op_type != "Gather":
291            logger.debug("failed to get gather")
292            return False
293
294        segment_initializer = self.get_initializer(gather.input[0])
295        if segment_initializer is None:
296            logger.debug("failed to get segment initializer")
297            return False
298
299        temp = numpy_helper.to_array(segment_initializer)
300        if len(temp.shape) == 2:
301            logger.info(f"Found segment embedding. name:{segment_initializer.name}, shape:{temp.shape}")
302            segment_embedding = segment_initializer.name
303        else:
304            logger.info(f"Failed to find segment embedding. name:{segment_initializer.name}, shape:{temp.shape}")
305            return False
306
307        logger.info("Create Embedding node")
308        self.create_embedding_subgraph(node, word_embedding, segment_embedding, position_embedding)
309        return True
310
311    def process_embedding(self):
312        """
313        Automatically detect word, segment and position embeddings.
314        """
315        logger.info("start processing embedding layer...")
316        output_name_to_node = self.output_name_to_node()
317        for node in self.nodes():
318            if node.op_type == "LayerNormalization":
319                if self.fuse_embedding(node, output_name_to_node):
320                    return
321                break
322
323    def fuse_mask(self):
324        nodes_to_remove = []
325        for node in self.nodes():
326            if node.op_type == "Mul" and self.has_constant_input(node, -10000):
327                mask_path = self.match_parent_path(node, ["Sub", "Cast", "Slice", "Unsqueeze"], [0, 1, 0, 0])
328                if mask_path is None:
329                    continue
330                sub_node, cast_node, slice_node, unsqueeze_node = mask_path
331
332                mask_input_name = self.attention_mask.get_first_mask()
333                if unsqueeze_node.input[0] != mask_input_name:
334                    print(f"Cast input {unsqueeze_node.input[0]} is not mask input {mask_input_name}")
335                    continue
336
337                unsqueeze_added_1 = onnx.helper.make_node(
338                    "Unsqueeze",
339                    inputs=[mask_input_name],
340                    outputs=["mask_fuse_unsqueeze1_output"],
341                    name="Mask_UnSqueeze_1",
342                    axes=[1],
343                )
344
345                unsqueeze_added_2 = onnx.helper.make_node(
346                    "Unsqueeze",
347                    inputs=["mask_fuse_unsqueeze1_output"],
348                    outputs=["mask_fuse_unsqueeze2_output"],
349                    name="Mask_UnSqueeze_2",
350                    axes=[2],
351                )
352
353                # self.replace_node_input(cast_node, cast_node.input[0], 'mask_fuse_unsqueeze2_output')
354                cast_node_2 = onnx.helper.make_node(
355                    "Cast",
356                    inputs=["mask_fuse_unsqueeze2_output"],
357                    outputs=["mask_fuse_cast_output"],
358                )
359                cast_node_2.attribute.extend([onnx.helper.make_attribute("to", 1)])
360                self.replace_node_input(sub_node, sub_node.input[1], "mask_fuse_cast_output")
361
362                nodes_to_remove.extend([slice_node, unsqueeze_node, cast_node])
363                self.add_node(unsqueeze_added_1)
364                self.add_node(unsqueeze_added_2)
365                self.add_node(cast_node_2)
366
367        self.remove_nodes(nodes_to_remove)
368
369        # Prune graph is done after removing nodes to remove island nodes.
370        if len(nodes_to_remove) > 0:
371            self.prune_graph()
372
373        logger.info("Fused mask" if len(nodes_to_remove) > 0 else "Failed to fuse mask")
374
375    def remove_extra_reshape(self):
376        skiplayernorm_nodes = self.get_nodes_by_op_type("SkipLayerNormalization")
377        reshape_removed = 0
378        for skiplayernorm_node in skiplayernorm_nodes:
379            path = self.match_parent_path(
380                skiplayernorm_node,
381                [
382                    "Add",
383                    "Reshape",
384                    "MatMul",
385                    "Reshape",
386                    "Gelu",
387                    "Add",
388                    "Reshape",
389                    "MatMul",
390                    "SkipLayerNormalization",
391                ],
392                [0, 0, 0, 0, 0, 0, 0, 0, 0],
393            )
394            if path is None:
395                continue
396
397            (
398                add_1,
399                reshape_1,
400                matmul_1,
401                reshape_2,
402                gelu,
403                add_2,
404                reshape_3,
405                matmul_2,
406                skiplayernorm,
407            ) = path
408            add_2.input[0] = matmul_2.output[0]
409            self.remove_node(reshape_3)
410            matmul_1.input[0] = gelu.output[0]
411            self.remove_node(reshape_2)
412            add_1.input[0] = matmul_1.output[0]
413            self.remove_node(reshape_1)
414            reshape_removed += 3
415
416        return reshape_removed
417
418    def remove_extra_reshape_2(self):
419        skiplayernorm_nodes = self.get_nodes_by_op_type("SkipLayerNormalization")
420        reshape_removed = 0
421        for skiplayernorm_node in skiplayernorm_nodes:
422            path = self.match_parent_path(
423                skiplayernorm_node,
424                [
425                    "Add",
426                    "Reshape",
427                    "MatMul",
428                    "Reshape",
429                    "Gelu",
430                    "Add",
431                    "Reshape",
432                    "MatMul",
433                    "Reshape",
434                    "SkipLayerNormalization",
435                ],
436                [None, 0, 0, 0, 0, 0, 0, 0, 0, 0],
437            )
438            if path is None:
439                continue
440
441            (
442                add_1,
443                reshape_1,
444                matmul_1,
445                reshape_2,
446                gelu,
447                add_2,
448                reshape_3,
449                matmul_2,
450                reshape_4,
451                skiplayernorm,
452            ) = path
453
454            matmul_2.input[0] = skiplayernorm.output[0]
455            self.remove_node(reshape_4)
456
457            add_2.input[0] = matmul_2.output[0]
458            self.remove_node(reshape_3)
459
460            matmul_1.input[0] = gelu.output[0]
461            self.remove_node(reshape_2)
462
463            add_1.input[0] = matmul_1.output[0]
464            self.remove_node(reshape_1)
465
466            reshape_removed += 4
467
468        return reshape_removed
469
470    def postprocess(self):
471        reshape_removed = self.remove_extra_reshape() + self.remove_extra_reshape_2()
472        logger.info(f"Remove {reshape_removed} Reshape nodes.")
473
474        self.prune_graph()
475 
codekingpro/portable-devtools · Team Ai