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 onnx import helper
9from onnx_model import OnnxModel
10
11logger = getLogger(__name__)
12
13
14class FusionBiasSplitGelu(Fusion):
15 def __init__(self, model: OnnxModel):
16 super().__init__(model, "BiasSplitGelu", "Gelu")
17
18 def fuse(self, gelu_node, input_name_to_nodes: dict, output_name_to_node: dict):
19 """
20 [root] --->Add --------------------> Slice ---------------> Mul -->
21 | ^ ^
22 | | |
23 +----------------------------+---Slice --> Gelu---+
24 | | ^
25 | |-----|
26 | | |
27 | Mul Mul
28 | ^ ^
29 v | |
30 Shape ---> Gather --> Add --> Div --+
31 """
32 if gelu_node.output[0] not in input_name_to_nodes:
33 return
34 children = input_name_to_nodes[gelu_node.output[0]]
35 if len(children) != 1 or children[0].op_type != "Mul":
36 return
37 mul_after_gelu = children[0]
38
39 slice_before_gelu = self.model.match_parent(gelu_node, "Slice", 0, output_name_to_node)
40 if slice_before_gelu is None:
41 return
42
43 if self.model.find_constant_input(slice_before_gelu, -1, delta=0.001) != 3:
44 return
45
46 add_output = slice_before_gelu.input[0]
47
48 start_index_nodes = self.model.match_parent_path(
49 slice_before_gelu,
50 ["Div", "Add", "Gather", "Shape", "Add"],
51 [1, 0, 0, 0, 0],
52 output_name_to_node, # Mul(1) is optional
53 )
54 if start_index_nodes is None:
55 start_index_nodes = self.model.match_parent_path(
56 slice_before_gelu,
57 ["Mul", "Div", "Add", "Gather", "Shape", "Add"],
58 [1, 0, 0, 0, 0, 0],
59 output_name_to_node,
60 )
61
62 if start_index_nodes is None or start_index_nodes[-2].input[0] != add_output:
63 return
64
65 end_index_nodes = self.model.match_parent_path(slice_before_gelu, ["Mul", "Div"], [2, 0], output_name_to_node)
66
67 if (
68 end_index_nodes is None or end_index_nodes[1] not in start_index_nodes
69 ): # the Div is parent of both two Mul nodes
70 return
71
72 slice_before_mul = self.model.match_parent(mul_after_gelu, "Slice", 0, output_name_to_node)
73 if slice_before_mul is None:
74 return
75
76 if (
77 slice_before_mul.input[2] != slice_before_gelu.input[1]
78 ): # end index of slice_before_mul is start index of slice_before_gelu
79 return
80
81 subgraph_nodes = [
82 *start_index_nodes,
83 end_index_nodes[0],
84 mul_after_gelu,
85 gelu_node,
86 slice_before_mul,
87 slice_before_gelu,
88 ]
89 subgraph_output = mul_after_gelu.output[0]
90 if not self.model.is_safe_to_fuse_nodes(
91 subgraph_nodes, [subgraph_output], input_name_to_nodes, output_name_to_node
92 ):
93 logger.info("Skip fuse BiasSplitGelu since it is not safe to fuse the subgraph.")
94 return
95
96 add_node = start_index_nodes[-1]
97 bias_index, _value = self.model.get_constant_input(add_node)
98 if not isinstance(bias_index, int):
99 return
100 self.nodes_to_remove.extend(subgraph_nodes)
101 node_name = self.model.create_node_name("BiasSplitGelu", name_prefix="BiasSplitGelu")
102 fused_node = helper.make_node(
103 "BiasSplitGelu",
104 inputs=[add_node.input[1 - bias_index], add_node.input[bias_index]],
105 outputs=[subgraph_output],
106 name=node_name,
107 )
108 fused_node.domain = "com.microsoft"
109 self.nodes_to_add.append(fused_node)
110 self.node_name_to_graph_name[node_name] = self.this_graph_name
111 