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 fusion_utils import NumpyHelper
9from onnx import helper
10from onnx_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionSkipGroupNorm(Fusion):
16 """
17 Fuse Add + GroupNorm into one node: SkipGroupNorm.
18 """
19
20 def __init__(self, model: OnnxModel):
21 super().__init__(model, "SkipGroupNorm", "GroupNorm")
22 # Update shape inference is needed since other fusions might add new edge which does not have shape info yet.
23 self.shape_infer_helper = self.model.infer_runtime_shape(update=True)
24
25 if self.shape_infer_helper is None:
26 logger.warning("SkipGroupNorm fusion will be skipped since symbolic shape inference disabled or failed.")
27
28 def create_transpose_node(self, input_name: str, perm: list[int], output_name=None):
29 """Append a Transpose node after an input"""
30 node_name = self.model.create_node_name("Transpose")
31 if output_name is None:
32 output_name = node_name + "_out" + "-" + input_name
33 transpose_node = helper.make_node("Transpose", inputs=[input_name], outputs=[output_name], name=node_name)
34 transpose_node.attribute.extend([helper.make_attribute("perm", perm)])
35 return transpose_node
36
37 def get_skip_index(self, add, is_channel_last: bool):
38 """Add has two inputs. This classifies which input is skip based on shape info (skip allows broadcast)."""
39 skip = -1
40 broadcast = False
41
42 assert self.shape_infer_helper is not None
43 shape_a = self.shape_infer_helper.get_edge_shape(add.input[0])
44 shape_b = self.shape_infer_helper.get_edge_shape(add.input[1])
45 assert shape_a is not None and shape_b is not None
46
47 if len(shape_a) == 4 and len(shape_b) == 4:
48 if shape_a == shape_b:
49 skip = 1
50 else:
51 c = 3 if is_channel_last else 1
52 h = 1 if is_channel_last else 2
53 w = 2 if is_channel_last else 3
54 if shape_a[0] == shape_b[0] and shape_a[c] == shape_b[c]:
55 if shape_b[h] == 1 and shape_b[w] == 1:
56 skip = 1
57 broadcast = True
58 elif shape_a[h] == 1 and shape_a[w] == 1:
59 skip = 0
60 broadcast = True
61
62 if skip < 0:
63 logger.debug(
64 "skip SkipGroupNorm fusion since shape of Add inputs (%s, %s) are not expected",
65 add.input[0],
66 add.input[1],
67 )
68 return skip, broadcast
69
70 def has_multiple_consumers(self, output_name, input_name_to_nodes):
71 """Whether an output has multiple consumers (like graph output or more than one children nodes)"""
72 return self.model.find_graph_output(output_name) is not None or (
73 output_name in input_name_to_nodes and len(input_name_to_nodes[output_name]) > 1
74 )
75
76 def remove_if_safe(self, node, input_name_to_nodes):
77 """Remove a node if it is safe (only one children, and not graph output)"""
78 if not self.has_multiple_consumers(node.output[0], input_name_to_nodes):
79 self.nodes_to_remove.extend([node])
80
81 def is_bias_1d(self, bias_name: str):
82 """Whether bias is an initializer of one dimension"""
83 initializer = self.model.get_initializer(bias_name)
84 if initializer is None:
85 return False
86
87 bias_weight = NumpyHelper.to_array(initializer)
88 if bias_weight is None:
89 logger.debug("Bias weight not found")
90 return False
91
92 if len(bias_weight.shape) != 1:
93 logger.debug("Bias weight is not 1D")
94 return False
95 return True
96
97 def match_bias_path(self, node, input_name_to_nodes, output_name_to_node):
98 """
99 Match the bias graph pattern from an Transpose node after Reshape node like in below example.
100 It checks whether the bias is 1D initializer. If so, remove Add and redirect MatMul output to Reshape.
101 """
102 # Before Fusion:
103 # MatMul (bias)
104 # \ / (shape)
105 # Add /
106 # \ /
107 # (a) Reshape
108 # \ |
109 # Transpose([0, 3, 1, 2]) Transpose([0, 3, 1, 2]) --- the start node, this func only handles the above nodes.
110 # \ /
111 # Add
112 # / \
113 # (c) Transpose([0,2,3,1])
114 # |
115 # GroupNorm
116 # |
117 # (d)
118 #
119 # After Fusion (the nodes below Reshape is handled in the fuse function):
120 # MatMul (shape)
121 # \ /
122 # (a) Reshape
123 # \ /
124 # SkipGroupNorm
125 # / \
126 # (d) Transpose([0, 3, 1, 2])
127 # \
128 # (c)
129
130 add_input_index = []
131 bias_nodes = self.model.match_parent_path(
132 node, ["Reshape", "Add", "MatMul"], [0, 0, None], output_name_to_node, add_input_index
133 )
134 if bias_nodes is None:
135 return None
136
137 (reshape, add_bias, matmul) = bias_nodes
138 bias = bias_nodes[1].input[1 - add_input_index[0]]
139 if not self.is_bias_1d(bias):
140 return None
141
142 reshape.input[0] = matmul.output[0]
143 self.remove_if_safe(add_bias, input_name_to_nodes)
144
145 return bias
146
147 def match_transpose_from_nhwc(self, output_name, input_name_to_nodes, output_name_to_node):
148 """Match whether an output is from a Transpose(perm=[0,3,1,2]) node."""
149 parent = output_name_to_node.get(output_name, None)
150 if parent is not None and parent.op_type == "Transpose":
151 permutation = OnnxModel.get_node_attribute(parent, "perm")
152 if permutation == [0, 3, 1, 2]:
153 self.remove_if_safe(parent, input_name_to_nodes)
154 return parent
155 return None
156
157 def fuse(self, node, input_name_to_nodes, output_name_to_node):
158 # This fusion requires shape information, so skip it if shape is not available.
159 if self.shape_infer_helper is None:
160 return
161
162 # Before Fusion:
163 # (a) (b)
164 # \ /
165 # Add
166 # /\
167 # (c) Transpose([0,2,3,1])
168 # \
169 # GroupNorm
170 # |
171 # (d)
172 #
173 # After Fusion:
174 # (a) (b)
175 # \ /
176 # Transpose([0,2,3,1]) Transpose([0,2,3,1])
177 # \ /
178 # SkipGroupNorm
179 # / \
180 # / Transpose([0, 3, 1, 2])
181 # / \
182 # (d) (c)
183 nodes = self.model.match_parent_path(node, ["Transpose", "Add"], [0, 0], output_name_to_node)
184 if nodes is None:
185 return
186
187 (transpose, add) = nodes
188 if transpose in self.nodes_to_remove or add in self.nodes_to_remove:
189 return
190
191 if self.has_multiple_consumers(transpose.output[0], input_name_to_nodes):
192 return
193
194 permutation = OnnxModel.get_node_attribute(transpose, "perm")
195 if permutation != [0, 2, 3, 1]:
196 return
197
198 inputs = []
199 bias = None
200 for i in range(2):
201 matched_transpose = self.match_transpose_from_nhwc(add.input[i], input_name_to_nodes, output_name_to_node)
202 if matched_transpose:
203 # When there is an Transpose node before Add (see examples in match_bias_path), we do not need to
204 # insert another Transpose node. The existing Transpose node will be removed in prune_graph if it
205 # has only one consumer.
206 inputs.append(matched_transpose.input[0])
207 # See whether it match bias pattern.
208 if bias is None:
209 bias = self.match_bias_path(matched_transpose, input_name_to_nodes, output_name_to_node)
210 else:
211 # Otherwise, insert a Transpose node before Add.
212 new_transpose = self.create_transpose_node(add.input[i], [0, 2, 3, 1])
213 self.model.add_node(new_transpose, self.this_graph_name)
214 inputs.append(new_transpose.output[0])
215
216 skip, broadcast = self.get_skip_index(add, is_channel_last=False)
217 if skip < 0:
218 return
219
220 inputs = [inputs[1 - skip], node.input[1], node.input[2], inputs[skip]]
221 if bias:
222 inputs = [*inputs, bias]
223
224 outputs = node.output
225
226 new_node_name = self.model.create_node_name(self.fused_op_type, name_prefix="SkipGroupNorm")
227 if self.has_multiple_consumers(add.output[0], input_name_to_nodes):
228 add_out_name = new_node_name + "_add_out"
229 outputs.append(add_out_name)
230
231 # Insert a Transpose node after add output.
232 add_out_transpose = self.create_transpose_node(add_out_name, [0, 3, 1, 2], add.output[0])
233 self.model.add_node(add_out_transpose, self.this_graph_name)
234
235 skip_group_norm = helper.make_node(
236 self.fused_op_type,
237 inputs=inputs,
238 outputs=outputs,
239 name=new_node_name,
240 )
241 skip_group_norm.domain = "com.microsoft"
242
243 self.increase_counter(
244 f"SkipGroupNorm(add_out={int(len(outputs) > 1)} bias={int(bias is not None)} broadcast={int(broadcast)})"
245 )
246
247 # Pass attributes from GroupNorm node to SkipGroupNorm
248 for att in node.attribute:
249 skip_group_norm.attribute.extend([att])
250
251 self.nodes_to_remove.extend([add, transpose, node])
252 self.nodes_to_add.append(skip_group_norm)
253 self.node_name_to_graph_name[skip_group_norm.name] = self.this_graph_name
254 self.prune_graph = True
255 