codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6from logging import getLogger
7
8from fusion_base import Fusion
9from fusion_utils import NumpyHelper
10from onnx import helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16def _is_broadcast_skip(input_shape, skip_shape):
17 """Check if skip_shape can broadcast to input_shape for SkipLayerNormalization.
18
19 The kernel supports: input 3D (B,S,H) with skip 3D (1,S,H) or skip 2D (S,H).
20 """
21 if len(input_shape) != 3:
22 return False
23 if len(skip_shape) == 3:
24 return skip_shape[0] == 1 and skip_shape[1] == input_shape[1] and skip_shape[2] == input_shape[2]
25 if len(skip_shape) == 2:
26 return skip_shape[0] == input_shape[1] and skip_shape[1] == input_shape[2]
27 return False
28
29
30class FusionSkipLayerNormalization(Fusion):
31 """
32 Fuse Add + LayerNormalization into one node: SkipLayerNormalization.
33 Supports broadcasting of the skip input: (1, sequence_length, hidden_size)
34 or (sequence_length, hidden_size) will be broadcast to match the input shape.
35 """
36
37 def __init__(
38 self,
39 model: OnnxModel,
40 fused_op_type: str = "SkipLayerNormalization",
41 search_op_types: str = "LayerNormalization",
42 shape_infer: bool = True,
43 ):
44 super().__init__(model, fused_op_type, search_op_types)
45 if shape_infer:
46 # Update shape inference is needed since other fusions might add new edge which does not have shape info yet.
47 self.shape_infer_helper = self.model.infer_runtime_shape({"batch_size": 4, "seq_len": 7}, update=True)
48 if self.shape_infer_helper is None:
49 # TODO(tianleiwu): support subgraph in shape inference.
50 logger.warning("symbolic shape inference disabled or failed.")
51
52 def get_skip_index(self, add):
53 """Identify which Add input is the skip tensor (the one that may broadcast).
54
55 Returns (skip_index, broadcast):
56 skip_index: 0 or 1 (which Add input is skip), -1 if incompatible
57 broadcast: True if broadcasting is needed
58 """
59 shape_a = self.shape_infer_helper.get_edge_shape(add.input[0])
60 shape_b = self.shape_infer_helper.get_edge_shape(add.input[1])
61 if shape_a is None or shape_b is None:
62 return -1, False
63
64 if shape_a == shape_b:
65 return (1, False) if len(shape_a) == 3 else (-1, False)
66
67 # Check if b is a broadcastable skip for a
68 if _is_broadcast_skip(shape_a, shape_b):
69 return 1, True
70 # Check if a is a broadcastable skip for b
71 if _is_broadcast_skip(shape_b, shape_a):
72 return 0, True
73
74 return -1, False
75
76 def fuse(self, node, input_name_to_nodes, output_name_to_node):
77 add = self.model.get_parent(node, 0, output_name_to_node)
78
79 # In some models there is input_ids->gather->add->LayerNorm and one of input of the
80 # add node is initializer with fixed shape which should not be fused into SkipLayerNorm
81 if add is None or add.op_type != "Add":
82 return
83
84 # The number of inputs of add should be 2
85 if len(add.input) != 2:
86 return
87
88 for add_input in add.input:
89 if self.model.get_initializer(add_input) is not None:
90 return
91
92 # To avoid an Add node have two children of LayerNormalization, we shall only fuse one SkipLayerNormalization
93 if add in self.nodes_to_remove:
94 return
95
96 # Root Mean Square Layer Normalization
97 simplified = node.op_type == "SimplifiedLayerNormalization"
98
99 skip_index = 1 # default: add.input[1] is the skip
100 _broadcast = False
101
102 if hasattr(self, "shape_infer_helper"):
103 if self.shape_infer_helper is not None:
104 skip_index, _broadcast = self.get_skip_index(add)
105 if skip_index < 0:
106 logger.debug(
107 "skip SkipLayerNormalization fusion since shapes of inputs (%s, %s) are not compatible",
108 add.input[0],
109 add.input[1],
110 )
111 return
112 else:
113 logger.debug("skip SkipLayerNormalization fusion since symbolic shape inference failed")
114 return
115
116 gather_path = self.model.match_parent_path(add, ["Gather"], [None])
117 if gather_path is not None and self.model.find_graph_input(gather_path[0].input[1]) is None:
118 if self.model.match_parent_path(gather_path[0], ["ConstantOfShape"], [1]) is None:
119 return
120
121 # When broadcasting is needed, check that neither Add input comes from a Gather
122 # (embedding lookup). Embedding Add+LayerNorm should be fused by EmbedLayerNormalization
123 # later in the pipeline, not as SkipLayerNormalization.
124 if _broadcast:
125 for i in range(2):
126 parent = self.model.get_parent(add, i, output_name_to_node)
127 if parent is not None and parent.op_type == "Gather":
128 logger.debug(
129 "skip SkipLayerNormalization broadcast fusion since Add input %d comes from Gather (embedding)",
130 i,
131 )
132 return
133
134 # This means that the residual Add before the LayerNormalization produces an output
135 # that is consumed by some other nodes or graph output other than the LayerNormalization itself
136 # We can still go ahead with the SkipLayerNormalization fusion but we need to
137 # preserve the output of Add and that needs to be produced by SkipLayerNormalization.
138 add_has_graph_output = self.model.find_graph_output(add.output[0]) is not None
139 residual_add_has_multiple_consumers = (
140 add_has_graph_output or len(self.model.get_children(add, input_name_to_nodes)) > 1
141 )
142
143 outputs_to_keep = node.output
144
145 if residual_add_has_multiple_consumers:
146 outputs_to_keep.extend([add.output[0]])
147
148 outputs = [node.output[0]]
149
150 # Skip the other optional outputs of SkipLayerNormalization before adding the Add's output
151 if residual_add_has_multiple_consumers:
152 outputs.extend(["", "", add.output[0]])
153
154 if self.model.is_safe_to_fuse_nodes([add, node], outputs_to_keep, input_name_to_nodes, output_name_to_node):
155 self.nodes_to_remove.extend([add, node])
156
157 input_index = 1 - skip_index
158 inputs = (
159 [add.input[input_index], add.input[skip_index], node.input[1], node.input[2]]
160 if not simplified
161 else [add.input[input_index], add.input[skip_index], node.input[1]]
162 )
163 normalize_node = helper.make_node(
164 self.fused_op_type,
165 inputs=inputs,
166 outputs=outputs,
167 name=self.model.create_node_name(self.fused_op_type, name_prefix="SkipLayerNorm"),
168 )
169 normalize_node.domain = "com.microsoft"
170
171 # Pass attribute "epsilon" from layernorm node to SkipLayerNormalization
172 for att in node.attribute:
173 if att.name == "epsilon":
174 normalize_node.attribute.extend([att])
175
176 # Set default epsilon if no epsilon exists from layernorm
177 if len(normalize_node.attribute) == 0:
178 normalize_node.attribute.extend([helper.make_attribute("epsilon", 1.0e-12)])
179
180 self.nodes_to_add.append(normalize_node)
181 self.node_name_to_graph_name[normalize_node.name] = self.this_graph_name
182
183
184class FusionBiasSkipLayerNormalization(Fusion):
185 def __init__(self, model: OnnxModel):
186 super().__init__(model, "SkipLayerNormalization", "SkipLayerNormalization", "add bias")
187
188 def fuse(self, node, input_name_to_nodes, output_name_to_node):
189 if len(node.input) != 4:
190 return
191
192 return_indice = []
193 nodes = self.model.match_parent_path(node, ["Add", "MatMul"], [None, None], output_name_to_node, return_indice)
194 if nodes is not None:
195 (add, _matmul) = nodes
196 else:
197 # In case of fp16, we could have a Cast between the MatMul and the bias Add
198 return_indice = []
199 nodes = self.model.match_parent_path(
200 node, ["Add", "Cast", "MatMul"], [None, None, None], output_name_to_node, return_indice
201 )
202 if nodes is not None:
203 (add, _cast, _matmul) = nodes
204 else:
205 return
206
207 assert len(return_indice) == 2 or len(return_indice) == 3
208 add_input_index = return_indice[0]
209 if add_input_index >= 2:
210 return
211 sln_input = add.input[return_indice[1]]
212 bias_input = add.input[1 - return_indice[1]]
213 skip_input = node.input[1 - add_input_index]
214
215 # bias should be one dimension
216 initializer = self.model.get_initializer(bias_input)
217 if initializer is None:
218 return
219 bias_weight = NumpyHelper.to_array(initializer)
220 if bias_weight is None:
221 logger.debug("Bias weight not found")
222 return
223 if len(bias_weight.shape) != 1:
224 logger.debug("Bias weight is not 1D")
225 return
226
227 subgraph_nodes = [node, add]
228 if not self.model.is_safe_to_fuse_nodes(subgraph_nodes, node.output, input_name_to_nodes, output_name_to_node):
229 logger.debug("Skip fusing SkipLayerNormalization with Bias since it is not safe")
230 return
231
232 self.nodes_to_remove.extend(subgraph_nodes)
233 inputs = [
234 sln_input,
235 skip_input,
236 node.input[2],
237 node.input[3],
238 bias_input,
239 ]
240 new_node = helper.make_node(
241 "SkipLayerNormalization",
242 inputs=inputs,
243 outputs=node.output,
244 name=self.model.create_node_name("SkipLayerNormalization", "SkipLayerNorm_AddBias_"),
245 )
246 new_node.domain = "com.microsoft"
247
248 # Pass attribute "epsilon" from skiplayernorm node to skiplayernorm(add bias)
249 for att in node.attribute:
250 if att.name == "epsilon":
251 new_node.attribute.extend([att])
252
253 # Set default epsilon if no epsilon exists from skiplayernorm
254 if len(new_node.attribute) == 0:
255 new_node.attribute.extend([helper.make_attribute("epsilon", 1.0e-12)])
256
257 self.nodes_to_add.append(new_node)
258 self.node_name_to_graph_name[new_node.name] = self.this_graph_name
259 