Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_gelu.py273 linesDownload Raw Back to fusions
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 
codekingpro/portable-devtools · Team Ai