codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5import logging
6
7import numpy as np
8from fusion_attention import AttentionMask, FusionAttention
9from onnx import helper
10from onnx_model import OnnxModel
11
12logger = logging.getLogger(__name__)
13
14
15class FusionBartAttention(FusionAttention):
16 """
17 Fuse Bart Attention subgraph into one Attention node.
18 """
19
20 def __init__(
21 self,
22 model: OnnxModel,
23 hidden_size: int,
24 num_heads: int,
25 attention_mask: AttentionMask,
26 ):
27 super().__init__(model, hidden_size, num_heads, attention_mask)
28
29 def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
30 # SkipLayerNormalization has two inputs, and one of them is the root input for attention.
31 qkv_nodes = self.model.match_parent_path(
32 normalize_node,
33 ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
34 [1, 1, 0, 0, 0],
35 )
36
37 # For LayerNormalization (when SkipLayerNorm fusion doesn't run, e.g. SDPA models where
38 # symbolic shape inference fails), there's an extra Add node for the residual connection
39 # between the LayerNorm and the attention output path.
40 add_before_layernorm = None
41 if qkv_nodes is None:
42 qkv_nodes_with_residual = self.model.match_parent_path(
43 normalize_node,
44 ["Add", "Add", "MatMul", "Reshape", "Transpose", "MatMul"],
45 [0, None, 0, 0, 0, 0],
46 )
47 if qkv_nodes_with_residual is not None:
48 add_before_layernorm = qkv_nodes_with_residual[0]
49 qkv_nodes = qkv_nodes_with_residual[1:]
50
51 if qkv_nodes is not None:
52 (
53 add_out,
54 matmul_out,
55 reshape_qkv,
56 transpose_qkv,
57 matmul_qkv,
58 ) = qkv_nodes
59 else:
60 logger.debug("fuse_attention: failed to match qkv path")
61 return
62
63 if add_before_layernorm is not None:
64 # LayerNorm case: root_input is the non-attention input of the residual Add
65 if add_before_layernorm.input[0] == add_out.output[0]:
66 root_input = add_before_layernorm.input[1]
67 else:
68 root_input = add_before_layernorm.input[0]
69 else:
70 other_inputs = []
71 for input_ in normalize_node.input:
72 if input_ not in output_name_to_node:
73 continue
74 if input_ == qkv_nodes[0].output[0]:
75 continue
76 other_inputs.append(input_)
77 if len(other_inputs) != 1:
78 return
79 root_input = other_inputs[0]
80
81 # Sometimes the input name to the attention MatMul nodes does not match the input name to the end
82 # SkipLayerNormalization node (name saved in root_input). We find the true input name to the MatMul
83 # nodes by getting the initial SkipLayerNormalization node and checking how many MatMul nodes are
84 # children nodes for each of its output names.
85 """
86 root_input
87 +---------------------------------------------------+
88 | |
89 | |
90 SkipLayerNormalization --> Attention --> MatMul --> SkipLayerNormalization
91 """
92 skip_layernorm = output_name_to_node[root_input]
93 # For some attention blocks, the end SkipLayerNormalization node may point to another node whose
94 # child is the LayerNormalization node.
95 if skip_layernorm.op_type in {"Add", "Clip"}:
96 skip_layernorm = self.model.get_children(skip_layernorm)[0]
97 for output in skip_layernorm.output:
98 if not output:
99 continue
100 children = input_name_to_nodes[output]
101 children_types = [child.op_type for child in children]
102 if children_types.count("MatMul") >= 1:
103 root_input = output
104 break
105
106 graph_input_names = {node.name for node in self.model.graph().input}
107 graph_output_names = {node.name for node in self.model.graph().output}
108
109 v_nodes_past_or_present = self.model.match_parent_path(
110 matmul_qkv,
111 ["Transpose", "Reshape", "Add", "MatMul"],
112 [1, 0, 0, None],
113 )
114 v_nodes_with_past = self.model.match_parent_path(
115 matmul_qkv,
116 ["Concat", "Transpose", "Reshape", "Add", "MatMul"],
117 [1, 1, 0, 0, None],
118 )
119 v_nodes_past_only_oai = self.model.match_parent_path(
120 matmul_qkv,
121 ["Transpose", "Reshape", "Reshape", "Transpose"],
122 [1, 0, 0, 0],
123 )
124 past_v, present_v = "", ""
125 v_nodes, add_v, matmul_v = [], None, None
126 if v_nodes_past_or_present is not None:
127 v_nodes = v_nodes_past_or_present
128 (transpose_v, reshape_v, add_v, matmul_v) = v_nodes
129
130 # Find past_v input name
131 start_child_nodes = input_name_to_nodes[add_v.output[0]]
132 for start_child_node in start_child_nodes:
133 if start_child_node.op_type == "Concat":
134 concat_v_nodes = self.model.match_parent_path(
135 start_child_node,
136 ["Reshape", "Transpose"],
137 [0, 0],
138 )
139 if concat_v_nodes is not None:
140 past_v = concat_v_nodes[-1].input[0]
141 start_child_nodes = input_name_to_nodes[start_child_node.output[0]]
142 break
143
144 # Find present_v output name
145 for start_child_node in start_child_nodes:
146 start_grandchild_nodes = input_name_to_nodes[start_child_node.output[0]]
147 for start_grandchild_node in start_grandchild_nodes:
148 if start_grandchild_node.output[0] in graph_output_names:
149 present_v = start_grandchild_node.output[0]
150 break
151 if present_v != "":
152 break
153 elif v_nodes_with_past is not None:
154 v_nodes = v_nodes_with_past
155 (concat_v, transpose_v, reshape_v, add_v, matmul_v) = v_nodes
156 past_v = concat_v.input[0]
157 present_v = concat_v.output[0]
158 elif matmul_qkv.input[1] in graph_input_names:
159 # Hugging Face's cross-attention where past_v is used directly as value
160 past_v = matmul_qkv.input[1]
161 elif v_nodes_past_only_oai is not None:
162 # OpenAI's cross-attention where past_v is used directly as value
163 v_nodes = v_nodes_past_only_oai
164 past_v = v_nodes[-1].input[0]
165 else:
166 logger.debug("fuse_attention: failed to match v path")
167 return
168 past_v = past_v if past_v in graph_input_names else ""
169 present_v = present_v if present_v in graph_output_names else ""
170
171 qk_nodes_no_mask = self.model.match_parent_path(matmul_qkv, ["Softmax", "MatMul"], [0, 0])
172 qk_nodes_with_mask = self.model.match_parent_path(matmul_qkv, ["Softmax", "Add", "MatMul"], [0, 0, 0])
173 # SDPA: NaN guard (Where(IsNaN, 0, softmax)) wraps the Softmax output.
174 # Where input[2] is the Softmax output (value when condition is False).
175 qk_nodes_sdpa_no_mask = self.model.match_parent_path(matmul_qkv, ["Where", "Softmax", "MatMul"], [0, 2, 0])
176 qk_nodes_sdpa_with_mask = self.model.match_parent_path(
177 matmul_qkv, ["Where", "Softmax", "Add", "MatMul"], [0, 2, 0, 0]
178 )
179 qk_nodes, add_qk = [], None
180 if qk_nodes_no_mask is not None:
181 _, matmul_qk = qk_nodes_no_mask
182 qk_nodes = qk_nodes_no_mask
183 elif qk_nodes_with_mask is not None:
184 _, add_qk, matmul_qk = qk_nodes_with_mask
185 qk_nodes = qk_nodes_with_mask
186 elif qk_nodes_sdpa_no_mask is not None:
187 _, _, matmul_qk = qk_nodes_sdpa_no_mask
188 qk_nodes = qk_nodes_sdpa_no_mask
189 elif qk_nodes_sdpa_with_mask is not None:
190 _, _, add_qk, matmul_qk = qk_nodes_sdpa_with_mask
191 qk_nodes = qk_nodes_sdpa_with_mask
192 else:
193 logger.debug("fuse_attention: failed to match qk path")
194 return
195
196 q_nodes_hf = self.model.match_parent_path(
197 matmul_qk,
198 ["Transpose", "Reshape", "Mul", "Add", "MatMul"],
199 [0, 0, 0, 0, 1],
200 )
201 q_nodes_oai = self.model.match_parent_path(
202 matmul_qk,
203 ["Mul", "Transpose", "Reshape", "Add", "MatMul"],
204 [0, 0, 0, 0, 1],
205 )
206 # SDPA: Mul(scale) applied before Transpose, MatMul may be at any Add input.
207 q_nodes_sdpa = self.model.match_parent_path(
208 matmul_qk,
209 ["Mul", "Transpose", "Reshape", "Add", "MatMul"],
210 [0, 0, 0, 0, None],
211 )
212 q_nodes = []
213 if q_nodes_hf is not None:
214 q_nodes = q_nodes_hf
215 (transpose_q, reshape_q, mul_q, add_q, matmul_q) = q_nodes
216 elif q_nodes_oai is not None:
217 q_nodes = q_nodes_oai
218 (mul_q, transpose_q, reshape_q, add_q, matmul_q) = q_nodes
219 elif q_nodes_sdpa is not None:
220 q_nodes = q_nodes_sdpa
221 (mul_q, transpose_q, reshape_q, add_q, matmul_q) = q_nodes
222 else:
223 logger.debug("fuse_attention: failed to match q path")
224 return
225
226 k_nodes_no_past_hf = self.model.match_parent_path(
227 matmul_qk,
228 ["Transpose", "Reshape", "MatMul"],
229 [1, 0, 0],
230 )
231 k_nodes_with_past_hf = self.model.match_parent_path(
232 matmul_qk,
233 ["Transpose", "Concat", "Transpose", "Reshape", "MatMul"],
234 [1, 0, 1, 0, 0],
235 )
236 k_nodes_past_or_present_oai = self.model.match_parent_path(
237 matmul_qk,
238 ["Mul", "Transpose", "Reshape", "MatMul"],
239 [1, 0, 0, 0],
240 )
241 k_nodes_past_only_oai = self.model.match_parent_path(
242 matmul_qk,
243 ["Mul", "Transpose", "Reshape", "Reshape", "Transpose"],
244 [1, 0, 0, 0, 0],
245 )
246 # SDPA: K is scaled (Mul) and transposed via Reshape->Transpose(0,2,1)->Reshape chain.
247 k_nodes_sdpa = self.model.match_parent_path(
248 matmul_qk,
249 ["Mul", "Reshape", "Transpose", "Reshape", "Transpose", "Reshape", "Add", "MatMul"],
250 [1, 0, 0, 0, 0, 0, 0, None],
251 )
252 past_k, present_k = "", ""
253 k_nodes, add_k, matmul_k = [], None, None
254 if k_nodes_no_past_hf is not None:
255 k_nodes = k_nodes_no_past_hf
256 (transpose_k, reshape_k, matmul_k) = k_nodes
257
258 # Find present_k output name
259 transpose_k_nodes = input_name_to_nodes[reshape_k.output[0]]
260 for transpose_k_node in transpose_k_nodes:
261 if transpose_k_node.output[0] in graph_output_names:
262 present_k = transpose_k_node.output[0]
263 break
264 elif k_nodes_with_past_hf is not None:
265 k_nodes = k_nodes_with_past_hf
266 (_, concat_k, transpose_k, reshape_k, matmul_k) = k_nodes
267 past_k = concat_k.input[0]
268 present_k = concat_k.output[0]
269 elif output_name_to_node[matmul_qk.input[1]].input[0] in graph_input_names:
270 # Hugging Face's cross-attention where past_k is used directly as key
271 k_nodes = [output_name_to_node[matmul_qk.input[1]]]
272 past_k = k_nodes[0].input[0]
273 elif k_nodes_sdpa is not None:
274 k_nodes = k_nodes_sdpa
275 (_, _, _, _, transpose_k, reshape_k, add_k, matmul_k) = k_nodes
276 elif k_nodes_past_or_present_oai is not None:
277 k_nodes = k_nodes_past_or_present_oai
278 (_, transpose_k, reshape_k, matmul_k) = k_nodes
279
280 # Find past_k input name
281 start_child_nodes = input_name_to_nodes[matmul_k.output[0]]
282 for start_child_node in start_child_nodes:
283 if start_child_node.op_type == "Concat":
284 concat_k_nodes = self.model.match_parent_path(
285 start_child_node,
286 ["Reshape", "Transpose"],
287 [0, 0],
288 )
289 if concat_k_nodes is not None:
290 past_k = concat_k_nodes[-1].input[0]
291 start_child_nodes = input_name_to_nodes[start_child_node.output[0]]
292 break
293
294 # Find present_k output name
295 for start_child_node in start_child_nodes:
296 start_grandchild_nodes = input_name_to_nodes[start_child_node.output[0]]
297 for start_grandchild_node in start_grandchild_nodes:
298 if start_grandchild_node.output[0] in graph_output_names:
299 present_k = start_grandchild_node.output[0]
300 break
301 if present_k != "":
302 break
303 elif k_nodes_past_only_oai is not None:
304 # OpenAI's cross-attention where past_k is used directly as key
305 k_nodes = k_nodes_past_only_oai
306 past_k = k_nodes[-1].input[0]
307 else:
308 logger.debug("fuse_attention: failed to match k path")
309 return
310 past_k = past_k if past_k in graph_input_names else ""
311 present_k = present_k if present_k in graph_output_names else ""
312
313 if matmul_k is not None and add_k is None:
314 # Create empty Add node for attention graph
315 add_v_tensor = self.model.get_initializer(add_v.input[0])
316 bias_dim = add_v_tensor.dims[0]
317 dtype = add_v_tensor.data_type
318 empty_bias_name = "empty_bias"
319 empty_tensor = self.model.get_initializer(empty_bias_name)
320 if empty_tensor is None:
321 self.add_initializer(
322 empty_bias_name,
323 dtype,
324 dims=[bias_dim],
325 vals=np.array([0.0] * bias_dim, dtype=helper.tensor_dtype_to_np_dtype(dtype)),
326 )
327
328 add_name = self.model.create_node_name("Add")
329 add_k = helper.make_node("Add", [empty_bias_name, matmul_k.output[0]], [reshape_k.name], add_name)
330
331 three_root_inputs = bool(past_k) and bool(past_v) and matmul_k is None and matmul_v is None
332 one_root_input = (
333 not three_root_inputs
334 and matmul_q.input[0] == root_input
335 and matmul_k.input[0] == root_input
336 and matmul_v.input[0] == root_input
337 )
338 two_root_inputs = (
339 not three_root_inputs
340 and matmul_q.input[0] == root_input
341 and matmul_k.input[0] == matmul_v.input[0]
342 and matmul_k.input[0] != matmul_q.input[0]
343 )
344
345 # There are 5 types of attention:
346 # 1) Encoder attention with one_root_input=True and no mask
347 # 2) Decoder self attention with one_root_input=True and has mask
348 # 3) Decoder cross attention with two_root_inputs=True and no mask
349 # 4) Decoder self attention with past with one_root_input=True and has mask and past_k and past_v
350 # 5) Decoder cross attention with past with three_root_inputs=True and no mask
351 # Derive mask presence from which QK pattern matched rather than re-walking the graph.
352 # This reuses the result of match_parent_paths above, which already tried both masked and
353 # unmasked variants and returned the first successful match.
354 has_mask = qk_nodes in (qk_nodes_with_mask, qk_nodes_sdpa_with_mask)
355 no_mask = not has_mask
356 encoder_attention = one_root_input and no_mask
357 decoder_self_attention = one_root_input and has_mask
358 decoder_cross_attention = two_root_inputs and no_mask
359 decoder_self_attention_with_past = decoder_self_attention and bool(past_k) and bool(past_v)
360 decoder_cross_attention_with_past = three_root_inputs and no_mask
361
362 # For decoder self-attentions, the attention mask needs to be included in the attention node
363 causal_mask = has_mask
364 mask_nodes = []
365 if causal_mask:
366 mask_nodes_bart = self.model.match_parent_path(
367 add_qk,
368 ["Where"],
369 [1],
370 )
371 mask_nodes_whisper_hf = self.model.match_parent_path(
372 add_qk,
373 ["Slice", "Expand", "Where"],
374 [1, 0, 1],
375 )
376 mask_nodes_whisper_oai = self.model.match_parent_path(
377 add_qk,
378 ["Slice", "Unsqueeze", "Gather", "Shape", "Add"],
379 [1, 2, 0, 0, 0],
380 )
381 mask_nodes_whisper_oai_unit_test = self.model.match_parent_path(
382 add_qk,
383 ["Slice", "Slice"],
384 [1, 0],
385 )
386 if mask_nodes_whisper_hf is not None:
387 mask_nodes = mask_nodes_whisper_hf
388 elif mask_nodes_whisper_oai is not None:
389 mask_nodes = mask_nodes_whisper_oai
390 elif mask_nodes_whisper_oai_unit_test is not None:
391 mask_nodes = mask_nodes_whisper_oai_unit_test
392 elif mask_nodes_bart is not None:
393 mask_nodes = mask_nodes_bart
394 else:
395 logger.debug("fuse_attention: failed to match mask nodes")
396 return
397 assert len(mask_nodes) > 0
398
399 if (
400 encoder_attention
401 or decoder_self_attention
402 or decoder_cross_attention
403 or decoder_self_attention_with_past
404 or decoder_cross_attention_with_past
405 ):
406 attention_last_node = reshape_qkv
407 num_heads, hidden_size = self.get_num_heads_and_hidden_size(reshape_q)
408
409 # Fall back to user-specified values when detected values are invalid
410 # (e.g., SDPA models use -1 in reshape shapes for dynamic dimensions).
411 if (num_heads <= 0 or hidden_size <= 0) and self.num_heads > 0 and self.hidden_size > 0:
412 logger.debug(
413 "fuse_attention: reshape dims invalid (num_heads=%d, hidden_size=%d), "
414 "falling back to user-specified num_heads=%d, hidden_size=%d",
415 num_heads,
416 hidden_size,
417 self.num_heads,
418 self.hidden_size,
419 )
420 num_heads = self.num_heads
421 hidden_size = self.hidden_size
422
423 if num_heads <= 0 or hidden_size <= 0 or (hidden_size % num_heads) != 0:
424 logger.debug("fuse_attention: failed to detect num_heads or hidden_size")
425 return
426
427 new_node = None
428 if decoder_self_attention_with_past or decoder_cross_attention or decoder_cross_attention_with_past:
429 # Note: Decoder attention with past key and past value is fused as multi-head attention
430 # rather than attention because multi-head attention supports separate past key and past
431 # value whereas attention supports concatenated past key and past value.
432 new_node = (
433 self.create_multihead_attention_node(
434 q_matmul=matmul_q,
435 k_matmul=matmul_k if decoder_cross_attention or decoder_self_attention_with_past else past_k,
436 v_matmul=matmul_v if decoder_cross_attention or decoder_self_attention_with_past else past_v,
437 q_add=add_q,
438 k_add=add_k if decoder_cross_attention or decoder_self_attention_with_past else None,
439 v_add=add_v if decoder_cross_attention or decoder_self_attention_with_past else None,
440 num_heads=num_heads,
441 hidden_size=hidden_size,
442 output=attention_last_node.output[0],
443 unidirectional=causal_mask,
444 past_k=past_k if decoder_self_attention_with_past else "",
445 past_v=past_v if decoder_self_attention_with_past else "",
446 present_k=present_k,
447 present_v=present_v,
448 )
449 if self.use_multi_head_attention
450 else None
451 )
452 else:
453 # Temporarily set multi-head attention flag to false
454 use_multi_head_attention_ground_truth = self.use_multi_head_attention
455 self.use_multi_head_attention = False
456 new_node = self.create_attention_node(
457 mask_index=None,
458 q_matmul=matmul_q,
459 k_matmul=matmul_k,
460 v_matmul=matmul_v,
461 q_add=add_q,
462 k_add=add_k,
463 v_add=add_v,
464 num_heads=num_heads,
465 hidden_size=hidden_size,
466 first_input=root_input,
467 output=attention_last_node.output[0],
468 causal=causal_mask,
469 past_k=past_k,
470 past_v=past_v,
471 present_k=present_k,
472 present_v=present_v,
473 )
474 self.use_multi_head_attention = use_multi_head_attention_ground_truth
475 if new_node is None:
476 logger.debug("fuse_attention: failed to create fused node")
477 return
478
479 self.nodes_to_add.append(new_node)
480 self.node_name_to_graph_name[new_node.name] = self.this_graph_name
481
482 self.nodes_to_remove.extend([attention_last_node, transpose_qkv, matmul_qkv])
483 self.nodes_to_remove.extend(qk_nodes)
484
485 # When using multi-head attention, keep MatMul nodes in original graph
486 if decoder_self_attention_with_past or decoder_cross_attention or decoder_cross_attention_with_past:
487 if len(q_nodes) > 0 and q_nodes[-1].op_type == "MatMul":
488 q_nodes.pop()
489 if len(k_nodes) > 0 and k_nodes[-1].op_type == "MatMul":
490 k_nodes.pop()
491 if len(v_nodes) > 0 and v_nodes[-1].op_type == "MatMul":
492 v_nodes.pop()
493 if self.disable_multi_head_attention_bias:
494 if len(q_nodes) > 0 and q_nodes[-1].op_type == "Add":
495 q_nodes.pop()
496 if len(k_nodes) > 0 and k_nodes[-1].op_type == "Add":
497 k_nodes.pop()
498 if len(v_nodes) > 0 and v_nodes[-1].op_type == "Add":
499 v_nodes.pop()
500
501 self.nodes_to_remove.extend(q_nodes)
502 self.nodes_to_remove.extend(k_nodes)
503 self.nodes_to_remove.extend(v_nodes)
504
505 # Use prune graph to remove mask nodes since they are shared by all attention nodes.
506 self.prune_graph = True
507 