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
7import numpy as np
8from fusion_base import Fusion
9from fusion_utils import FusionUtils
10from onnx import helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionGptAttentionPastBase(Fusion):
17 """Base class for GPT Attention Fusion with past state"""
18
19 def __init__(self, model: OnnxModel, num_heads: int):
20 super().__init__(model, "Attention", ["LayerNormalization", "SkipLayerNormalization"], "with past")
21 self.num_heads = num_heads
22 self.utils = FusionUtils(model)
23 self.casted_attention_mask = {} # map from name of attention mask to the name that casted to int32
24 self.mask_filter_value = None
25
26 def match_past_pattern_1(self, concat_k, concat_v, output_name_to_node):
27 # Pattern 1:
28 # {past}
29 # / \
30 # / \
31 # Gather(axes=0, indices=0) Gather(indices=1)
32 # | |
33 # Transpose (perm=0,1,3,2) |
34 # | |
35 # Concat_k Concat_v
36 # | /
37 # Transpose (perm=0,1,3,2) /
38 # | /
39 # Unsqueeze Unsqueeze
40 # \ /
41 # \ /
42 # Concat
43 # |
44 # {present}
45 gather = self.model.get_parent(concat_v, 0, output_name_to_node)
46 if gather is None or gather.op_type != "Gather":
47 logger.debug("match_past_pattern_1: expect Gather for past")
48 return None
49
50 if self.model.find_constant_input(gather, 1) != 1:
51 logger.debug("match_past_pattern_1: expect indices=1 for Gather of past")
52 return None
53 past = gather.input[0]
54
55 parent = self.model.get_parent(concat_k, 0, output_name_to_node)
56 if parent and parent.op_type == "Gather":
57 gather_past_k = parent
58 else:
59 past_k_nodes = self.model.match_parent_path(concat_k, ["Transpose", "Gather"], [0, 0])
60 if past_k_nodes is None:
61 logger.debug("match_past_pattern_1: failed match Transpose and Gather")
62 return None
63 gather_past_k = past_k_nodes[-1]
64
65 if self.model.find_constant_input(gather_past_k, 0) != 1:
66 logger.debug("match_past_pattern_1: expect indices=0 for Gather k of past")
67 return None
68 past_k = gather_past_k.input[0]
69 if past != past_k:
70 logger.debug("match_past_pattern_1: expect past to be same")
71 return None
72
73 return past
74
75 def match_past_pattern_2(self, concat_k, concat_v, output_name_to_node):
76 # Pattern 2:
77 # Split (QKV)
78 # / | |
79 # / | +----------------------+
80 # | |
81 # | {past} |
82 # | | |
83 # Reshape Split Reshape
84 # | / \ |
85 # Transpose_k Squeeze Squeeze Transpose_v
86 # | | \ /
87 # +------|---+ \ /
88 # | | \ /
89 # Concat_k Concat_v
90 # | |
91 # Unsqueeze Unsqueeze
92 # \ /
93 # Concat
94 # |
95 # {present}
96 #
97 squeeze = self.model.get_parent(concat_v, 0, output_name_to_node)
98 if squeeze is None or squeeze.op_type != "Squeeze":
99 logger.debug("match_past_pattern_2: expect Squeeze as parent of concat_v")
100 return None
101
102 split = self.model.get_parent(squeeze, 0, output_name_to_node)
103 if split is None or split.op_type != "Split":
104 logger.debug("match_past_pattern_2: expect Split for past path")
105 return None
106
107 opset_version = self.model.get_opset_version()
108 if opset_version < 13:
109 if not FusionUtils.check_node_attribute(squeeze, "axes", [0]):
110 logger.debug("match_past_pattern_2: axes != [0] for Squeeze in past path")
111 return None
112
113 if not FusionUtils.check_node_attribute(split, "split", [1, 1]):
114 logger.debug("match_past_pattern_2: split != [1, 1] for Split in past path")
115 return None
116 else:
117 if not self.utils.check_node_input_value(squeeze, 1, [0]):
118 logger.debug("match_past_pattern_2: axes != [0] for Squeeze in past path")
119 return None
120
121 if not self.utils.check_node_input_value(split, 1, [1, 1]):
122 logger.debug("match_past_pattern_2: split != [1, 1] for Split in past path")
123 return None
124
125 if not FusionUtils.check_node_attribute(split, "axis", 0, default_value=0):
126 logger.debug("match_past_pattern_2: attribute axis of Split are not expected in past path")
127 return None
128 past = split.input[0]
129
130 past_k_nodes = self.model.match_parent_path(concat_k, ["Squeeze", "Split"], [0, 0])
131 if past_k_nodes is None:
132 logger.debug("match_past_pattern_2: failed to match past_k_nodes path")
133 return None
134 past_k = past_k_nodes[-1].input[0]
135
136 if past != past_k:
137 logger.info("match_past_pattern_2: expect past to be same")
138 return None
139
140 return past
141
142 def match_present(self, concat_v, input_name_to_nodes):
143 unsqueeze_present_v = self.model.find_first_child_by_type(
144 concat_v, "Unsqueeze", input_name_to_nodes, recursive=False
145 )
146 if not unsqueeze_present_v:
147 logger.info("expect unsqueeze for present")
148 return None
149 concat_present = self.model.find_first_child_by_type(
150 unsqueeze_present_v, "Concat", input_name_to_nodes, recursive=False
151 )
152 if not concat_present:
153 logger.info("expect concat for present")
154 return None
155
156 present = concat_present.output[0]
157 return present
158
159 def cast_attention_mask(self, input_name):
160 if input_name in self.casted_attention_mask:
161 attention_mask_input_name = self.casted_attention_mask[input_name]
162 elif self.model.find_graph_input(input_name):
163 casted, attention_mask_input_name = self.utils.cast_graph_input_to_int32(input_name)
164 self.casted_attention_mask[input_name] = attention_mask_input_name
165 else:
166 attention_mask_input_name, cast_node = self.utils.cast_input_to_int32(input_name)
167 self.casted_attention_mask[input_name] = attention_mask_input_name
168 return attention_mask_input_name
169
170
171class FusionGptAttention(FusionGptAttentionPastBase):
172 """
173 Fuse GPT-2 Attention with past state subgraph into one Attention node.
174 """
175
176 def __init__(self, model: OnnxModel, num_heads: int):
177 super().__init__(model, num_heads)
178
179 def create_attention_node(
180 self,
181 fc_weight,
182 fc_bias,
183 gemm_qkv,
184 past,
185 present,
186 input,
187 output,
188 mask,
189 is_unidirectional,
190 ):
191 attention_node_name = self.model.create_node_name("GptAttention")
192 attention_node = helper.make_node(
193 "Attention",
194 inputs=[input, fc_weight, fc_bias, mask, past],
195 outputs=[attention_node_name + "_output", present],
196 name=attention_node_name,
197 )
198 attention_node.domain = "com.microsoft"
199 attention_node.attribute.extend(
200 [
201 helper.make_attribute("num_heads", self.num_heads),
202 helper.make_attribute("unidirectional", 1 if is_unidirectional else 0),
203 ]
204 )
205
206 if self.mask_filter_value is not None:
207 attention_node.attribute.extend([helper.make_attribute("mask_filter_value", float(self.mask_filter_value))])
208
209 matmul_node = helper.make_node(
210 "MatMul",
211 inputs=[attention_node_name + "_output", gemm_qkv.input[1]],
212 outputs=[attention_node_name + "_matmul_output"],
213 name=attention_node_name + "_matmul",
214 )
215
216 add_node = helper.make_node(
217 "Add",
218 inputs=[attention_node_name + "_matmul_output", gemm_qkv.input[2]],
219 outputs=[output],
220 name=attention_node_name + "_add",
221 )
222 self.nodes_to_add.extend([attention_node, matmul_node, add_node])
223 self.node_name_to_graph_name[attention_node.name] = self.this_graph_name
224 self.node_name_to_graph_name[matmul_node.name] = self.this_graph_name
225 self.node_name_to_graph_name[add_node.name] = self.this_graph_name
226
227 def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
228 past = None
229 present = None
230 return_indice = []
231
232 is_normalize_node_skiplayernorm = normalize_node.op_type == "SkipLayerNormalization"
233 qkv_nodes = None
234
235 if not is_normalize_node_skiplayernorm:
236 qkv_nodes = self.model.match_parent_path(
237 normalize_node,
238 ["Add", "Reshape", "Gemm", "Reshape", "Reshape", "Transpose", "MatMul"],
239 [0, None, 0, 0, 0, 0, 0],
240 output_name_to_node=output_name_to_node,
241 return_indice=return_indice,
242 )
243 else:
244 qkv_nodes = self.model.match_parent_path(
245 normalize_node,
246 ["Reshape", "Gemm", "Reshape", "Reshape", "Transpose", "MatMul"],
247 [None, 0, 0, 0, 0, 0],
248 output_name_to_node=output_name_to_node,
249 return_indice=return_indice,
250 )
251
252 if qkv_nodes is None:
253 return
254
255 another_input = None
256 if not is_normalize_node_skiplayernorm:
257 (
258 add_qkv,
259 reshape_qkv,
260 gemm_qkv,
261 reshape_1,
262 reshape_2,
263 transpose_qkv,
264 matmul_qkv,
265 ) = qkv_nodes
266
267 another_input = add_qkv.input[1 - return_indice[0]]
268 else:
269 (
270 reshape_qkv,
271 gemm_qkv,
272 reshape_1,
273 reshape_2,
274 transpose_qkv,
275 matmul_qkv,
276 ) = qkv_nodes
277
278 v_nodes = self.model.match_parent_path(matmul_qkv, ["Concat", "Transpose", "Reshape", "Split"], [1, 1, 0, 0])
279 if v_nodes is None:
280 logger.debug("fuse_attention: failed to match v path")
281 return
282 (concat_v, transpose_v, reshape_v, split_fc) = v_nodes
283
284 # Try match pattern using Gemm + LayerNormalization
285 fc_nodes = self.model.match_parent_path(
286 split_fc,
287 ["Reshape", "Gemm", "Reshape", "LayerNormalization"],
288 [0, 0, 0, 0],
289 output_name_to_node,
290 )
291
292 # Try match pattern using Gemm + SkipLayerNormalization
293 if fc_nodes is None:
294 fc_nodes = self.model.match_parent_path(
295 split_fc,
296 ["Reshape", "Gemm", "Reshape", "SkipLayerNormalization"],
297 [0, 0, 0, 0],
298 output_name_to_node,
299 )
300
301 # Try match pattern using MatMul
302 if fc_nodes is None:
303 # LayerNormalization
304 fc_nodes = self.model.match_parent_path(
305 split_fc,
306 ["Add", "MatMul", "LayerNormalization"],
307 [0, None, 0],
308 output_name_to_node,
309 )
310
311 # SkipLayerNormalization
312 if fc_nodes is None:
313 fc_nodes = self.model.match_parent_path(
314 split_fc,
315 ["Add", "MatMul", "SkipLayerNormalization"],
316 [0, None, 0],
317 output_name_to_node,
318 )
319
320 if fc_nodes is None:
321 logger.debug("fuse_attention: failed to match fc path")
322 return
323
324 fc_weight = fc_nodes[1].input[1]
325 i, _ = self.model.get_constant_input(fc_nodes[0])
326 fc_bias = fc_nodes[0].input[i]
327 else:
328 fc_weight = fc_nodes[1].input[1]
329 fc_bias = fc_nodes[1].input[2]
330
331 layernorm_before_attention = fc_nodes[-1]
332
333 # `another_input` will be non-None only if
334 # (1) SkipLayerNorm fusion wasn't turned ON
335 # (2) SkipLayerNorm fusion was turned ON but upstream layer's LayerNorm + Add was not
336 # fused into a SkipLayerNorm. This can happen if the shapes to the Add node are different.
337 # So, keep the following check if SkipLayerNorm fusion is turned ON or OFF.
338 if another_input is not None and another_input not in layernorm_before_attention.input:
339 logger.debug("Upstream Add and (Skip)LayerNormalization shall have one same input")
340 return
341
342 is_unidirectional = True
343 slice_mask = None
344 input_mask_nodes = None
345 concat_k_to_match = None
346 qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "Sub", "Mul", "Div", "MatMul"], [0, 0, 0, 0, 0])
347 if qk_nodes is not None:
348 (softmax_qk, sub_qk, mul_qk, div_qk, matmul_qk) = qk_nodes
349 mask_nodes = self.model.match_parent_path(
350 sub_qk,
351 [
352 "Mul",
353 "Sub",
354 "Slice",
355 "Slice",
356 "Unsqueeze",
357 "Sub",
358 "Squeeze",
359 "Slice",
360 "Shape",
361 "Div",
362 ],
363 [1, 0, 1, 0, 1, 0, 0, 0, 0, 0],
364 )
365 if mask_nodes is None:
366 logger.debug("fuse_attention: failed to match unidirectional mask path")
367 return
368 div_mask = mask_nodes[-1]
369 slice_mask = mask_nodes[3]
370
371 if div_qk != div_mask:
372 logger.debug("fuse_attention: skip since div_qk != div_mask")
373 return
374
375 if len(mask_nodes) > 1 and mask_nodes[0].op_type == "Mul":
376 _, mul_val = self.model.get_constant_input(mask_nodes[0])
377 if mul_val != -10000:
378 self.mask_filter_value = -mul_val
379
380 else:
381 # New pattern for gpt2 from PyTorch 1.5.0 and Transformers 2.9.0.
382 i, qk_nodes, _ = self.model.match_parent_paths(
383 matmul_qkv,
384 [
385 (["Softmax", "Where", "Div", "MatMul"], [0, 0, 1, 0]),
386 (["Softmax", "Add", "Where", "Div", "MatMul"], [0, 0, None, 1, 0]),
387 ],
388 output_name_to_node,
389 )
390 if qk_nodes is None:
391 logger.debug("fuse_attention: failed to match qk nodes")
392 return
393
394 where_qk = qk_nodes[-3]
395 div_qk = qk_nodes[-2]
396 matmul_qk = qk_nodes[-1]
397
398 if i == 1:
399 add_qk = qk_nodes[1]
400 _, input_mask_nodes, _ = self.model.match_parent_paths(
401 add_qk,
402 [
403 (
404 ["Mul", "Sub", "Cast", "Unsqueeze", "Unsqueeze", "Reshape"],
405 [None, 0, 1, 0, 0, 0],
406 ),
407 (
408 ["Mul", "Sub", "Unsqueeze", "Unsqueeze", "Reshape"],
409 [None, 0, 1, 0, 0],
410 ),
411 (
412 ["Mul", "Sub", "Unsqueeze", "Unsqueeze"],
413 [None, 0, 1, 0],
414 ), # useless cast and reshape are removed.
415 ],
416 output_name_to_node,
417 )
418 if input_mask_nodes is None:
419 logger.debug("fuse_attention: failed to match input attention mask path")
420 return
421 if len(input_mask_nodes) > 1 and input_mask_nodes[0].op_type == "Mul":
422 _, mul_val = self.model.get_constant_input(input_mask_nodes[0])
423 if mul_val != -10000:
424 self.mask_filter_value = mul_val
425
426 i, mask_nodes, _ = self.model.match_parent_paths(
427 where_qk,
428 [
429 (
430 ["Cast", "Slice", "Slice", "Unsqueeze", "Sub", "Squeeze", "Slice", "Shape"],
431 [0, 0, 0, 1, 0, 0, 0, 0],
432 ),
433 # For Transformers >= 4.27, causal mask uses torch.bool instead of torch.uint8, so no Cast to bool.
434 (
435 ["Slice", "Slice", "Unsqueeze", "Sub", "Squeeze", "Slice", "Shape"],
436 [0, 0, 1, 0, 0, 0, 0],
437 ),
438 ],
439 output_name_to_node,
440 )
441 if mask_nodes is None:
442 # TODO: match mask path for GPT2LMHeadModel_BeamSearchStep.
443 logger.debug("fuse_attention: failed to match mask path")
444 return
445
446 slice_mask = mask_nodes[2 if i == 0 else 1]
447
448 div_or_concat = self.model.get_parent(mask_nodes[-1], 0, output_name_to_node)
449 if div_or_concat.op_type == "Div":
450 div_mask = div_or_concat
451 if div_qk != div_mask:
452 logger.debug("fuse_attention: skip since div_qk != div_mask")
453 return
454 elif div_or_concat.op_type == "Concat":
455 concat_k_to_match = div_or_concat
456 else:
457 logger.debug("fuse_attention: failed to match mask path")
458
459 # Validate that the mask data is either lower triangular (unidirectional) or all ones
460 mask_data = self.model.get_constant_value(slice_mask.input[0])
461 if not (
462 isinstance(mask_data, np.ndarray)
463 and len(mask_data.shape) == 4
464 and mask_data.shape[:2] == (1, 1)
465 and mask_data.shape[2] == mask_data.shape[3]
466 ):
467 logger.debug("fuse_attention: skip since mask shape is not 1x1xWxW")
468 return
469
470 if np.allclose(mask_data, np.ones_like(mask_data)):
471 is_unidirectional = False
472 elif not np.allclose(mask_data, np.tril(np.ones_like(mask_data))):
473 logger.debug("fuse_attention: skip since mask is neither lower triangular nor ones")
474 return
475
476 q_nodes = self.model.match_parent_path(matmul_qk, ["Transpose", "Reshape", "Split"], [0, 0, 0])
477 if q_nodes is None:
478 logger.debug("fuse_attention: failed to match q path")
479 return
480 (transpose_q, reshape_q, split_q) = q_nodes
481 if split_fc != split_q:
482 logger.debug("fuse_attention: skip since split_fc != split_q")
483 return
484
485 k_nodes = self.model.match_parent_path(matmul_qk, ["Concat", "Transpose", "Reshape", "Split"], [1, 1, 0, 0])
486 if k_nodes is None:
487 # This pattern is from pytorch 1.7.1 and transformers 4.6.1
488 k_nodes = self.model.match_parent_path(
489 matmul_qk,
490 ["Transpose", "Concat", "Transpose", "Reshape", "Split"],
491 [1, 0, 1, 0, 0],
492 )
493 if k_nodes is None:
494 logger.debug("fuse_attention: failed to match k path")
495 return
496 else:
497 (_, concat_k, transpose_k, reshape_k, split_k) = k_nodes
498 else:
499 (concat_k, transpose_k, reshape_k, split_k) = k_nodes
500 if split_fc != split_k:
501 logger.debug("fuse_attention: skip since split_fc != split_k")
502 return
503
504 if concat_k_to_match and concat_k != concat_k_to_match:
505 logger.debug("fuse_attention: skip since concat_k != concat_k_to_match")
506 return
507
508 attention_mask_input_name = ""
509 if input_mask_nodes is not None:
510 input_name = input_mask_nodes[-1].input[0]
511 attention_mask_input_name = self.cast_attention_mask(input_name)
512
513 # Match past and present paths
514 past = self.match_past_pattern_1(concat_k, concat_v, output_name_to_node) or self.match_past_pattern_2(
515 concat_k, concat_v, output_name_to_node
516 )
517 if past is None:
518 logger.info("fuse_attention: failed to match past path")
519 return
520 if not self.model.find_graph_input(past):
521 logger.debug("past is not graph input.")
522 # For GPT2LMHeadModel_BeamSearchStep, there is an extra Gather node to select beam index so it is not graph input.
523
524 present = self.match_present(concat_v, input_name_to_nodes)
525 if present is None:
526 logger.info("fuse_attention: failed to match present path")
527 return
528 if not self.model.find_graph_output(present):
529 logger.info("expect present to be graph output")
530 return
531
532 self.create_attention_node(
533 fc_weight,
534 fc_bias,
535 gemm_qkv,
536 past,
537 present,
538 layernorm_before_attention.output[0],
539 reshape_qkv.output[0],
540 attention_mask_input_name,
541 is_unidirectional,
542 )
543
544 # we rely on prune_graph() to clean old subgraph nodes:
545 # qk_nodes + q_nodes + k_nodes + v_nodes + mask_nodes + [reshape_qkv, transpose_qkv, matmul_qkv]
546 self.prune_graph = True
547 