codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5from logging import getLogger
6
7from fusion_base import Fusion
8from numpy import ndarray
9from onnx import helper
10from onnx_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionBiasAdd(Fusion):
16 def __init__(self, model: OnnxModel):
17 super().__init__(model, "BiasAdd", "Add")
18
19 def fuse(self, add_node, input_name_to_nodes: dict, output_name_to_node: dict):
20 """
21 Fuse Add bias and Add skip connection into BiasAdd
22 """
23
24 nodes = self.model.match_parent_path(
25 add_node,
26 ["Add", "MatMul", "BiasSplitGelu", "MatMul", "SkipLayerNormalization"],
27 [0, None, 0, 0, 0],
28 output_name_to_node,
29 )
30
31 if nodes is None:
32 return
33
34 bias_node = nodes[0]
35 skip_layer_norm = nodes[-1]
36
37 # Check skip connection is from SkipLayerNormalization output
38 if add_node.input[1] not in skip_layer_norm.output:
39 return
40
41 bias_index, bias_value = self.model.get_constant_input(bias_node)
42 if not (isinstance(bias_index, int) and (bias_value is not None) and isinstance(bias_value, ndarray)):
43 return
44 if bias_value.ndim != 1:
45 return
46
47 self.nodes_to_remove.extend([add_node, bias_node])
48 node_name = self.model.create_node_name("BiasAdd")
49 fused_node = helper.make_node(
50 "BiasAdd",
51 inputs=[bias_node.input[1 - bias_index], bias_node.input[bias_index], add_node.input[1]],
52 outputs=[add_node.output[0]],
53 name=node_name,
54 )
55 fused_node.domain = "com.microsoft"
56 self.nodes_to_add.append(fused_node)
57 self.node_name_to_graph_name[node_name] = self.this_graph_name
58 