codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License. See License.txt in the project root for
4# license information.
5# --------------------------------------------------------------------------
6from __future__ import annotations
7
8import onnx
9
10from ..onnx_model import ONNXModel
11from .fusion import Fusion
12
13
14class FusionGelu(Fusion):
15 def __init__(self, model: ONNXModel):
16 super().__init__(model, "Gelu", "Erf")
17
18 def fuse(
19 self,
20 erf_node: onnx.NodeProto,
21 input_name_to_nodes: dict[str, list[onnx.NodeProto]],
22 output_name_to_node: dict[str, onnx.NodeProto],
23 ):
24 """
25 Interface function that tries to fuse a node sequence containing an Erf node into a single
26 Gelu node.
27 """
28 if (
29 self.fuse_1(erf_node, input_name_to_nodes, output_name_to_node)
30 or self.fuse_2(erf_node, input_name_to_nodes, output_name_to_node)
31 or self.fuse_3(erf_node, input_name_to_nodes, output_name_to_node)
32 ):
33 self.model.set_opset_import("com.microsoft", 1)
34
35 def fuse_1(
36 self,
37 erf_node: onnx.NodeProto,
38 input_name_to_nodes: dict[str, list[onnx.NodeProto]],
39 output_name_to_node: dict[str, onnx.NodeProto],
40 ) -> bool:
41 """
42 This pattern is from PyTorch model
43 Fuse Gelu with Erf into one node:
44 Pattern 1:
45 +-------Mul(0.5)---------------------+
46 | |
47 | v
48 [root] --> Div -----> Erf --> Add --> Mul -->
49 (B=1.4142...) (1)
50
51 Pattern 2:
52 +------------------------------------+
53 | |
54 | v
55 [root] --> Div -----> Erf --> Add --> Mul -->Mul -->
56 (B=1.4142...) (1) (0.5)
57
58 Note that constant input for Add and Mul could be first or second input: like either A=0.5 or B=0.5 is fine.
59 """
60 if erf_node.output[0] not in input_name_to_nodes:
61 return False
62 children = input_name_to_nodes[erf_node.output[0]]
63 if len(children) != 1 or children[0].op_type != "Add":
64 return False
65 add_after_erf = children[0]
66
67 if not self.has_constant_input(add_after_erf, 1):
68 return False
69
70 if add_after_erf.output[0] not in input_name_to_nodes:
71 return False
72
73 children = input_name_to_nodes[add_after_erf.output[0]]
74 if len(children) != 1 or children[0].op_type != "Mul":
75 return False
76
77 mul_after_erf = children[0]
78
79 div = self.match_parent(erf_node, "Div", 0, output_name_to_node)
80 if div is None:
81 return False
82
83 if self.find_constant_input(div, 1.4142, delta=0.001) != 1:
84 return False
85
86 subgraph_input = div.input[0]
87
88 another = 1 if mul_after_erf.input[0] == add_after_erf.output[0] else 0
89 if subgraph_input == mul_after_erf.input[another]: # pattern 2
90 children = input_name_to_nodes[mul_after_erf.output[0]]
91 if len(children) != 1 or children[0].op_type != "Mul":
92 return False
93 mul_half = children[0]
94 if not self.has_constant_input(mul_half, 0.5):
95 return False
96 subgraph_output = mul_half.output[0]
97 else: # pattern 1
98 mul_half = self.match_parent(mul_after_erf, "Mul", another, output_name_to_node)
99 if mul_half is None:
100 return False
101
102 if not self.has_constant_input(mul_half, 0.5):
103 return False
104
105 if subgraph_input not in mul_half.input:
106 return False
107
108 subgraph_output = mul_after_erf.output[0]
109
110 subgraph_nodes = [div, erf_node, add_after_erf, mul_after_erf, mul_half]
111 if not self.is_safe_to_fuse_nodes(subgraph_nodes, [subgraph_output], input_name_to_nodes, output_name_to_node):
112 return False
113
114 self.nodes_to_remove.extend(subgraph_nodes)
115 fused_node = onnx.helper.make_node(
116 "Gelu", name=self.create_unique_node_name(), inputs=[subgraph_input], outputs=[subgraph_output]
117 )
118 fused_node.domain = "com.microsoft"
119 self.nodes_to_add.append(fused_node)
120 return True
121
122 def fuse_2(
123 self,
124 erf_node: onnx.NodeProto,
125 input_name_to_nodes: dict[str, list[onnx.NodeProto]],
126 output_name_to_node: dict[str, onnx.NodeProto],
127 ) -> bool:
128 """
129 This pattern is from Keras model
130 Fuse Gelu with Erf into one node:
131 +------------------------------------------+
132 | |
133 | v
134 [root] --> Div -----> Erf --> Add --> Mul -->Mul
135 (B=1.4142...) (A=1) (A=0.5)
136
137 Note that constant input for Add and Mul could be first or second input: like either A=0.5 or B=0.5 is fine.
138 """
139 if erf_node.output[0] not in input_name_to_nodes:
140 return False
141 children = input_name_to_nodes[erf_node.output[0]]
142 if len(children) != 1 or children[0].op_type != "Add":
143 return False
144 add_after_erf = children[0]
145
146 if not self.has_constant_input(add_after_erf, 1):
147 return False
148
149 if add_after_erf.output[0] not in input_name_to_nodes:
150 return False
151 children = input_name_to_nodes[add_after_erf.output[0]]
152 if len(children) != 1 or children[0].op_type != "Mul":
153 return False
154 mul_after_erf = children[0]
155
156 if not self.has_constant_input(mul_after_erf, 0.5):
157 return False
158
159 if mul_after_erf.output[0] not in input_name_to_nodes:
160 return False
161 children = input_name_to_nodes[mul_after_erf.output[0]]
162 if len(children) != 1 or children[0].op_type != "Mul":
163 return False
164 mul = children[0]
165
166 div = self.match_parent(erf_node, "Div", 0, output_name_to_node)
167 if div is None:
168 return False
169
170 sqrt_node = None
171 if self.find_constant_input(div, 1.4142, delta=0.001) != 1:
172 sqrt_node = self.match_parent(div, "Sqrt", 1, output_name_to_node)
173 if sqrt_node is None:
174 return False
175 if not self.has_constant_input(sqrt_node, 2.0):
176 return False
177
178 subgraph_input = div.input[0]
179
180 if subgraph_input not in mul.input:
181 return False
182
183 subgraph_nodes = [div, erf_node, add_after_erf, mul_after_erf, mul]
184 if sqrt_node:
185 subgraph_nodes.append(sqrt_node)
186
187 if not self.is_safe_to_fuse_nodes(subgraph_nodes, [mul.output[0]], input_name_to_nodes, output_name_to_node):
188 return False
189
190 self.nodes_to_remove.extend(subgraph_nodes)
191 fused_node = onnx.helper.make_node(
192 "Gelu", name=self.create_unique_node_name(), inputs=[subgraph_input], outputs=[mul.output[0]]
193 )
194 fused_node.domain = "com.microsoft"
195 self.nodes_to_add.append(fused_node)
196 return True
197
198 def fuse_3(
199 self,
200 erf_node: onnx.NodeProto,
201 input_name_to_nodes: dict[str, list[onnx.NodeProto]],
202 output_name_to_node: dict[str, onnx.NodeProto],
203 ) -> bool:
204 """
205 This pattern is from TensorFlow model
206 Fuse Gelu with Erf into one node:
207 +----------------------------------------------+
208 | |
209 | v
210 [root] --> Mul -----> Erf --> Add --> Mul -->Mul
211 (A=0.7071067690849304) (B=1) (B=0.5)
212
213 Note that constant input for Add and Mul could be first or second input: like either A=0.5 or B=0.5 is fine.
214 """
215
216 if erf_node.output[0] not in input_name_to_nodes:
217 return False
218 children = input_name_to_nodes[erf_node.output[0]]
219 if len(children) != 1 or children[0].op_type != "Add":
220 return False
221 add_after_erf = children[0]
222
223 if not self.has_constant_input(add_after_erf, 1):
224 return False
225
226 if add_after_erf.output[0] not in input_name_to_nodes:
227 return False
228 children = input_name_to_nodes[add_after_erf.output[0]]
229 if len(children) != 1 or children[0].op_type != "Mul":
230 return False
231 mul_half = children[0]
232
233 if not self.has_constant_input(mul_half, 0.5):
234 return False
235
236 first_mul = self.match_parent(erf_node, "Mul", 0, output_name_to_node)
237 if first_mul is None:
238 return False
239
240 i = self.find_constant_input(first_mul, 0.7071067690849304, delta=0.001)
241 if i < 0:
242 return False
243
244 root_input_index = 1 - i
245 subgraph_input = first_mul.input[root_input_index]
246
247 if mul_half.output[0] not in input_name_to_nodes:
248 return False
249 children = input_name_to_nodes[mul_half.output[0]]
250 if len(children) != 1 or children[0].op_type != "Mul":
251 return False
252 last_mul = children[0]
253
254 if not (last_mul.input[0] == subgraph_input or last_mul.input[1] == subgraph_input):
255 return False
256
257 subgraph_nodes = [first_mul, erf_node, add_after_erf, mul_half, last_mul]
258 if not self.is_safe_to_fuse_nodes(
259 subgraph_nodes,
260 [last_mul.output[0]],
261 input_name_to_nodes,
262 output_name_to_node,
263 ):
264 return False
265
266 self.nodes_to_remove.extend(subgraph_nodes)
267 fused_node = onnx.helper.make_node(
268 "Gelu", name=self.create_unique_node_name(), inputs=[subgraph_input], outputs=[last_mul.output[0]]
269 )
270 fused_node.domain = "com.microsoft"
271 self.nodes_to_add.append(fused_node)
272 return True
273 