codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6import logging
7
8import onnx
9from onnx import numpy_helper
10from onnx_model_bert_tf import BertOnnxModelTF
11
12logger = logging.getLogger(__name__)
13
14
15class BertOnnxModelKeras(BertOnnxModelTF):
16 def __init__(self, model, num_heads, hidden_size):
17 super().__init__(model, num_heads, hidden_size)
18
19 def match_mask_path(self, add_or_sub_before_softmax):
20 mask_nodes = self.match_parent_path(
21 add_or_sub_before_softmax,
22 ["Mul", "Sub", "Reshape", "Cast"],
23 [1, None, 1, 0],
24 )
25 if mask_nodes is not None:
26 return mask_nodes
27
28 mask_nodes = self.match_parent_path(
29 add_or_sub_before_softmax,
30 ["Mul", "Sub", "Cast", "Slice", "Unsqueeze"],
31 [1, 1, 1, 0, 0],
32 )
33 if mask_nodes is not None:
34 return mask_nodes
35
36 mask_nodes = self.match_parent_path(
37 add_or_sub_before_softmax,
38 ["Mul", "Sub", "Cast", "Unsqueeze", "Unsqueeze"],
39 [1, None, 1, 0, 0],
40 )
41 return mask_nodes
42
43 def check_attention_input(self, matmul_q, matmul_k, matmul_v, parent, output_name_to_node):
44 reshape_nodes = []
45
46 for x in [matmul_q, matmul_k, matmul_v]:
47 root_input = x.input[0]
48 root_node = output_name_to_node[root_input]
49 if root_node == parent:
50 continue
51 if root_node.op_type == "Reshape" and root_node.input[0] == parent.output[0]:
52 reshape_nodes.append(root_node)
53 continue
54 logger.debug(f"Check attention input failed:{root_input}, {parent.output[0]}")
55 return False, []
56
57 return True, reshape_nodes
58
59 def fuse_attention(self):
60 self.input_name_to_nodes()
61 output_name_to_node = self.output_name_to_node()
62
63 nodes_to_remove = []
64 attention_count = 0
65
66 skip_layer_norm_nodes = self.get_nodes_by_op_type("SkipLayerNormalization")
67 for normalize_node in skip_layer_norm_nodes:
68 # SkipLayerNormalization has two inputs, and one of them is the root input for attention.
69 parent = self.get_parent(normalize_node, 0)
70 if parent is None or parent.op_type not in [
71 "SkipLayerNormalization",
72 "EmbedLayerNormalization",
73 ]:
74 if parent.op_type == "Add":
75 parent = self.get_parent(normalize_node, 1)
76 if parent is None or parent.op_type not in [
77 "SkipLayerNormalization",
78 "EmbedLayerNormalization",
79 ]:
80 logger.debug(f"First input for skiplayernorm: {parent.op_type if parent is not None else None}")
81 continue
82 else:
83 logger.debug(f"First input for skiplayernorm: {parent.op_type if parent is not None else None}")
84 continue
85 else:
86 # TODO: shall we add back the checking of children op types.
87 pass
88
89 qkv_nodes = self.match_parent_path(
90 normalize_node,
91 ["Add", "Reshape", "MatMul", "Reshape", "Transpose", "MatMul"],
92 [None, 0, 0, 0, 0, 0],
93 )
94 if qkv_nodes is None:
95 logger.debug("Failed to match qkv nodes")
96 continue
97 (
98 add,
99 extra_reshape_0,
100 matmul,
101 reshape_qkv,
102 transpose_qkv,
103 matmul_qkv,
104 ) = qkv_nodes
105 logger.debug("Matched qkv nodes")
106
107 v_nodes = self.match_parent_path(
108 matmul_qkv,
109 ["Transpose", "Reshape", "Add", "Reshape", "MatMul"],
110 [1, 0, 0, 0, 0],
111 )
112 if v_nodes is None:
113 logger.debug("Failed to match v path")
114 continue
115 (transpose_v, reshape_v, add_v, extra_reshape_1, matmul_v) = v_nodes
116
117 qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Sub", "MatMul"], [0, 0, 0])
118 if qk_nodes is not None:
119 (softmax_qk, sub_qk, matmul_qk) = qk_nodes
120 q_nodes = self.match_parent_path(
121 matmul_qk,
122 ["Mul", "Transpose", "Reshape", "Add", "Reshape", "MatMul"],
123 [0, None, 0, 0, 0, 0],
124 )
125 if q_nodes is not None:
126 (
127 mul_q,
128 transpose_q,
129 reshape_q,
130 add_q,
131 extra_reshape_2,
132 matmul_q,
133 ) = q_nodes
134
135 else:
136 qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Add", "Mul", "MatMul"], [0, 0, 0, None])
137 if qk_nodes is None:
138 qk_nodes = self.match_parent_path(matmul_qkv, ["Softmax", "Add", "Div", "MatMul"], [0, 0, 0, None])
139 if qk_nodes is None:
140 logger.debug("Failed to match qk path")
141 continue
142 (softmax_qk, add_qk, mul_qk, matmul_qk) = qk_nodes
143
144 q_nodes = self.match_parent_path(
145 matmul_qk,
146 ["Transpose", "Reshape", "Add", "Reshape", "MatMul"],
147 [0, 0, 0, 0, 0],
148 )
149 if q_nodes is not None:
150 (transpose_q, reshape_q, add_q, extra_reshape_2, matmul_q) = q_nodes
151
152 if q_nodes is None:
153 logger.debug("Failed to match q path")
154 continue
155
156 k_nodes = self.match_parent_path(
157 matmul_qk,
158 ["Transpose", "Reshape", "Add", "Reshape", "MatMul"],
159 [1, 0, 0, 0, 0],
160 )
161 if k_nodes is None:
162 logger.debug("Failed to match k path")
163 continue
164 (transpose_k, reshape_k, add_k, extra_reshape_3, matmul_k) = k_nodes
165
166 mask_nodes = self.match_mask_path(qk_nodes[1])
167 if mask_nodes is None:
168 logger.debug("Failed to match mask path")
169 continue
170 if not self.has_constant_input(mask_nodes[1], 1):
171 logger.debug("Sub node expected to have an input with constant value 1.0.")
172 continue
173
174 is_same_root, reshape_nodes = self.check_attention_input(
175 matmul_q, matmul_k, matmul_v, parent, output_name_to_node
176 )
177 if is_same_root:
178 mask_index = self.attention_mask.process_mask(mask_nodes[-1].input[0])
179 logger.debug("Create an Attention node.")
180 attention_node = self.attention_fusion.create_attention_node(
181 mask_index=mask_index,
182 q_matmul=matmul_q,
183 k_matmul=matmul_k,
184 v_matmul=matmul_v,
185 q_add=add_q,
186 k_add=add_k,
187 v_add=add_v,
188 num_heads=self.num_heads,
189 hidden_size=self.hidden_size,
190 first_input=parent.output[0],
191 output=reshape_qkv.output[0],
192 )
193 if attention_node is None:
194 continue
195
196 self.add_node(attention_node)
197 attention_count += 1
198
199 nodes_to_remove.extend([reshape_qkv, transpose_qkv, matmul_qkv])
200 nodes_to_remove.extend(qk_nodes)
201 nodes_to_remove.extend(q_nodes)
202 nodes_to_remove.extend(k_nodes)
203 nodes_to_remove.extend(v_nodes)
204 nodes_to_remove.extend(mask_nodes)
205 nodes_to_remove.extend(reshape_nodes)
206 nodes_to_remove.append(extra_reshape_0)
207 self.replace_node_input(add, extra_reshape_0.output[0], matmul.output[0])
208 else:
209 logger.debug("Root node not matched.")
210 continue
211 self.remove_nodes(nodes_to_remove)
212 self.update_graph()
213 logger.info(f"Fused Attention count:{attention_count}")
214
215 def preprocess(self):
216 self.process_embedding()
217 self.fuse_mask()
218 self.skip_reshape()
219
220 def skip_reshape(self):
221 self.input_name_to_nodes()
222 self.output_name_to_node()
223
224 count = 0
225 reshape_nodes = self.get_nodes_by_op_type("Reshape")
226 for reshape_node in reshape_nodes:
227 parent = self.get_parent(reshape_node, 0)
228 if parent is not None and parent.op_type == "Reshape":
229 reshape_node.input[0] = parent.input[0]
230 count += 1
231
232 if count > 0:
233 logger.info(f"Skip consequent Reshape count: {count}")
234
235 def fuse_embedding(self, node, output_name_to_node):
236 assert node.op_type == "LayerNormalization"
237 logger.debug(f"start fusing embedding from node with output={node.output[0]}...")
238 word_embed_path = self.match_parent_path(node, ["Add", "Add", "Gather"], [0, 0, 0], output_name_to_node)
239 if word_embed_path is None:
240 logger.debug("failed to match word_embed_path")
241 return False
242
243 skip_node, add_node, gather_node = word_embed_path
244
245 word_initializer = self.get_initializer(gather_node.input[0])
246 if word_initializer is None:
247 logger.debug("failed to get word initializer")
248 return False
249
250 temp = numpy_helper.to_array(word_initializer)
251 if len(temp.shape) == 2:
252 logger.info(f"Found word embedding. name:{word_initializer.name}, shape:{temp.shape}")
253 word_embedding = word_initializer.name
254 else:
255 logger.info(f"Failed to find word embedding. name:{word_initializer.name}, shape:{temp.shape}")
256 return False
257
258 pos_initializer = self.get_initializer(add_node.input[1])
259 if pos_initializer is not None:
260 temp = numpy_helper.to_array(pos_initializer)
261 if len(temp.shape) == 3 and temp.shape[0] == 1:
262 tensor = numpy_helper.from_array(temp.reshape((temp.shape[1], temp.shape[2])), "position_embedding")
263 self.add_initializer(tensor)
264 logger.info(f"Found position embedding. name:{pos_initializer.name}, shape:{temp.shape[1:]}")
265 position_embedding = "position_embedding"
266 else:
267 logger.info(f"Failed to find position embedding. name:{pos_initializer.name}, shape:{temp.shape}")
268 return False
269 else:
270 pos_embed_path = self.match_parent_path(add_node, ["Gather", "Slice"], [1, 1], output_name_to_node)
271 if pos_embed_path is None:
272 logger.debug("failed to match pos_embed_path")
273 return False
274
275 pos_gather, pos_slice = pos_embed_path
276 pos_initializer = self.get_initializer(pos_gather.input[0])
277 if pos_initializer is None:
278 logger.debug("failed to get pos initializer")
279 return False
280
281 temp = numpy_helper.to_array(pos_initializer)
282 if len(temp.shape) == 2:
283 logger.info(f"Found word embedding. name:{pos_initializer.name}, shape:{temp.shape}")
284 position_embedding = pos_initializer.name
285 else:
286 logger.info(f"Failed to find position embedding. name:{pos_initializer.name}, shape:{temp.shape}")
287 return False
288
289 gather = self.get_parent(skip_node, 1, output_name_to_node)
290 if gather is None or gather.op_type != "Gather":
291 logger.debug("failed to get gather")
292 return False
293
294 segment_initializer = self.get_initializer(gather.input[0])
295 if segment_initializer is None:
296 logger.debug("failed to get segment initializer")
297 return False
298
299 temp = numpy_helper.to_array(segment_initializer)
300 if len(temp.shape) == 2:
301 logger.info(f"Found segment embedding. name:{segment_initializer.name}, shape:{temp.shape}")
302 segment_embedding = segment_initializer.name
303 else:
304 logger.info(f"Failed to find segment embedding. name:{segment_initializer.name}, shape:{temp.shape}")
305 return False
306
307 logger.info("Create Embedding node")
308 self.create_embedding_subgraph(node, word_embedding, segment_embedding, position_embedding)
309 return True
310
311 def process_embedding(self):
312 """
313 Automatically detect word, segment and position embeddings.
314 """
315 logger.info("start processing embedding layer...")
316 output_name_to_node = self.output_name_to_node()
317 for node in self.nodes():
318 if node.op_type == "LayerNormalization":
319 if self.fuse_embedding(node, output_name_to_node):
320 return
321 break
322
323 def fuse_mask(self):
324 nodes_to_remove = []
325 for node in self.nodes():
326 if node.op_type == "Mul" and self.has_constant_input(node, -10000):
327 mask_path = self.match_parent_path(node, ["Sub", "Cast", "Slice", "Unsqueeze"], [0, 1, 0, 0])
328 if mask_path is None:
329 continue
330 sub_node, cast_node, slice_node, unsqueeze_node = mask_path
331
332 mask_input_name = self.attention_mask.get_first_mask()
333 if unsqueeze_node.input[0] != mask_input_name:
334 print(f"Cast input {unsqueeze_node.input[0]} is not mask input {mask_input_name}")
335 continue
336
337 unsqueeze_added_1 = onnx.helper.make_node(
338 "Unsqueeze",
339 inputs=[mask_input_name],
340 outputs=["mask_fuse_unsqueeze1_output"],
341 name="Mask_UnSqueeze_1",
342 axes=[1],
343 )
344
345 unsqueeze_added_2 = onnx.helper.make_node(
346 "Unsqueeze",
347 inputs=["mask_fuse_unsqueeze1_output"],
348 outputs=["mask_fuse_unsqueeze2_output"],
349 name="Mask_UnSqueeze_2",
350 axes=[2],
351 )
352
353 # self.replace_node_input(cast_node, cast_node.input[0], 'mask_fuse_unsqueeze2_output')
354 cast_node_2 = onnx.helper.make_node(
355 "Cast",
356 inputs=["mask_fuse_unsqueeze2_output"],
357 outputs=["mask_fuse_cast_output"],
358 )
359 cast_node_2.attribute.extend([onnx.helper.make_attribute("to", 1)])
360 self.replace_node_input(sub_node, sub_node.input[1], "mask_fuse_cast_output")
361
362 nodes_to_remove.extend([slice_node, unsqueeze_node, cast_node])
363 self.add_node(unsqueeze_added_1)
364 self.add_node(unsqueeze_added_2)
365 self.add_node(cast_node_2)
366
367 self.remove_nodes(nodes_to_remove)
368
369 # Prune graph is done after removing nodes to remove island nodes.
370 if len(nodes_to_remove) > 0:
371 self.prune_graph()
372
373 logger.info("Fused mask" if len(nodes_to_remove) > 0 else "Failed to fuse mask")
374
375 def remove_extra_reshape(self):
376 skiplayernorm_nodes = self.get_nodes_by_op_type("SkipLayerNormalization")
377 reshape_removed = 0
378 for skiplayernorm_node in skiplayernorm_nodes:
379 path = self.match_parent_path(
380 skiplayernorm_node,
381 [
382 "Add",
383 "Reshape",
384 "MatMul",
385 "Reshape",
386 "Gelu",
387 "Add",
388 "Reshape",
389 "MatMul",
390 "SkipLayerNormalization",
391 ],
392 [0, 0, 0, 0, 0, 0, 0, 0, 0],
393 )
394 if path is None:
395 continue
396
397 (
398 add_1,
399 reshape_1,
400 matmul_1,
401 reshape_2,
402 gelu,
403 add_2,
404 reshape_3,
405 matmul_2,
406 skiplayernorm,
407 ) = path
408 add_2.input[0] = matmul_2.output[0]
409 self.remove_node(reshape_3)
410 matmul_1.input[0] = gelu.output[0]
411 self.remove_node(reshape_2)
412 add_1.input[0] = matmul_1.output[0]
413 self.remove_node(reshape_1)
414 reshape_removed += 3
415
416 return reshape_removed
417
418 def remove_extra_reshape_2(self):
419 skiplayernorm_nodes = self.get_nodes_by_op_type("SkipLayerNormalization")
420 reshape_removed = 0
421 for skiplayernorm_node in skiplayernorm_nodes:
422 path = self.match_parent_path(
423 skiplayernorm_node,
424 [
425 "Add",
426 "Reshape",
427 "MatMul",
428 "Reshape",
429 "Gelu",
430 "Add",
431 "Reshape",
432 "MatMul",
433 "Reshape",
434 "SkipLayerNormalization",
435 ],
436 [None, 0, 0, 0, 0, 0, 0, 0, 0, 0],
437 )
438 if path is None:
439 continue
440
441 (
442 add_1,
443 reshape_1,
444 matmul_1,
445 reshape_2,
446 gelu,
447 add_2,
448 reshape_3,
449 matmul_2,
450 reshape_4,
451 skiplayernorm,
452 ) = path
453
454 matmul_2.input[0] = skiplayernorm.output[0]
455 self.remove_node(reshape_4)
456
457 add_2.input[0] = matmul_2.output[0]
458 self.remove_node(reshape_3)
459
460 matmul_1.input[0] = gelu.output[0]
461 self.remove_node(reshape_2)
462
463 add_1.input[0] = matmul_1.output[0]
464 self.remove_node(reshape_1)
465
466 reshape_removed += 4
467
468 return reshape_removed
469
470 def postprocess(self):
471 reshape_removed = self.remove_extra_reshape() + self.remove_extra_reshape_2()
472 logger.info(f"Remove {reshape_removed} Reshape nodes.")
473
474 self.prune_graph()
475 