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_options import AttentionMaskFormat
10from fusion_utils import FusionUtils, NumpyHelper
11from onnx import NodeProto, TensorProto, helper, numpy_helper
12from onnx_model import OnnxModel
13
14logger = getLogger(__name__)
15
16
17class AttentionMask:
18 """
19 Fuse Attention subgraph into one Attention node.
20 """
21
22 def __init__(self, model: OnnxModel):
23 self.model = model
24 # A lookup table with mask input as key, and mask index output as value
25 self.mask_indice = {}
26 # A lookup table with mask input as key, and cast (to int32) output as value
27 self.mask_casted = {}
28 self.utils = FusionUtils(model)
29 self.mask_format = AttentionMaskFormat.MaskIndexEnd
30 self.opset_version = model.get_opset_version()
31
32 def set_mask_format(self, mask_format: AttentionMaskFormat):
33 self.mask_format = mask_format
34
35 def set_mask_indice(self, mask, mask_index):
36 if mask in self.mask_indice:
37 assert mask_index == self.mask_indice[mask]
38 self.mask_indice[mask] = mask_index
39
40 def get_first_mask(self):
41 assert len(self.mask_indice) > 0
42 return next(iter(self.mask_indice))
43
44 def process_mask(self, mask_2d: str) -> str | None:
45 if self.mask_format == AttentionMaskFormat.NoMask:
46 return None
47
48 if mask_2d in self.mask_indice:
49 return self.mask_indice[mask_2d]
50
51 # Add cast to convert int64 to int32
52 if self.model.find_graph_input(mask_2d):
53 casted, input_name = self.utils.cast_graph_input_to_int32(mask_2d)
54 else:
55 input_name, _cast_node = self.utils.cast_input_to_int32(mask_2d)
56 casted = True
57
58 if casted:
59 self.mask_casted[mask_2d] = input_name
60
61 # Attention supports int32 attention mask (2D) since 1.4.0
62 if self.mask_format == AttentionMaskFormat.AttentionMask:
63 self.mask_indice[mask_2d] = input_name
64 return input_name
65
66 # Add a mask processing node to convert attention mask to mask index (1D)
67 output_name = self.model.create_node_name("mask_index")
68 if self.opset_version < 13:
69 mask_index_node = helper.make_node(
70 "ReduceSum",
71 inputs=[input_name],
72 outputs=[output_name],
73 name=self.model.create_node_name("ReduceSum", "MaskReduceSum"),
74 )
75 mask_index_node.attribute.extend([helper.make_attribute("axes", [1]), helper.make_attribute("keepdims", 0)])
76 else:
77 # ReduceSum-13: axes is moved from attribute to input
78 axes_name = "ort_const_1_reduce_sum_axes"
79 if self.model.get_initializer(axes_name) is None:
80 self.model.add_initializer(
81 helper.make_tensor(
82 name=axes_name,
83 data_type=TensorProto.INT64,
84 dims=[1],
85 vals=[1],
86 raw=False,
87 )
88 )
89 mask_index_node = helper.make_node(
90 "ReduceSum",
91 inputs=[input_name, axes_name],
92 outputs=[output_name],
93 name=self.model.create_node_name("ReduceSum", "MaskReduceSum"),
94 )
95 mask_index_node.attribute.extend([helper.make_attribute("keepdims", 0)])
96
97 self.model.add_node(mask_index_node)
98
99 self.mask_indice[mask_2d] = output_name
100 return output_name
101
102
103class FusionAttention(Fusion):
104 """
105 Fuse Attention subgraph into one Attention node.
106 """
107
108 def __init__(
109 self,
110 model: OnnxModel,
111 hidden_size: int,
112 num_heads: int,
113 attention_mask: AttentionMask | None = None,
114 use_multi_head_attention: bool = False,
115 disable_multi_head_attention_bias: bool = False,
116 search_op_types: list[str] = ["SkipLayerNormalization", "LayerNormalization"], # noqa: B006
117 ):
118 attention_op_name = "MultiHeadAttention" if use_multi_head_attention else "Attention"
119 super().__init__(model, attention_op_name, search_op_types)
120 self.hidden_size = hidden_size
121 self.num_heads = num_heads
122 self.attention_mask = attention_mask if attention_mask else AttentionMask(model)
123 self.use_multi_head_attention = use_multi_head_attention
124 self.disable_multi_head_attention_bias = disable_multi_head_attention_bias
125 self.mask_filter_value = None
126
127 # Flags to show warning only once
128 self.num_heads_warning = True
129 self.hidden_size_warning = True
130
131 self.shape_infer = None
132 self.shape_infer_done = True
133
134 def get_num_heads_and_hidden_size_from_concat(self, concat: NodeProto) -> tuple[int, int]:
135 """
136 Detect num_heads and hidden_size from Concat node in the following subgraph:
137
138 SkipLayerNormalization or EmbedLayerNormalization
139 / |
140 MatMul Shape
141 | |
142 Add Gather(indices=0)
143 | |
144 | Unsqueeze
145 | |
146 | Concat (*, -1, 12, 64)
147 | /
148 Reshape
149 |
150 Transpose
151 """
152 if len(concat.input) == 4:
153 num_heads = self.model.get_constant_value(concat.input[2])
154 head_size = self.model.get_constant_value(concat.input[3])
155 if (
156 isinstance(num_heads, np.ndarray)
157 and num_heads.size == 1
158 and isinstance(head_size, np.ndarray)
159 and head_size.size == 1
160 ):
161 return num_heads[0], num_heads[0] * head_size[0]
162
163 return self.num_heads, self.hidden_size
164
165 def get_num_heads_and_hidden_size(self, reshape_q: NodeProto) -> tuple[int, int]:
166 """Detect num_heads and hidden_size from a reshape node.
167
168 Args:
169 reshape_q (NodeProto): reshape node for Q
170
171 Returns:
172 Tuple[int, int]: num_heads and hidden_size
173 """
174 # we assume that reshape fusion has done, so the shape is a tensor like [0, 0, num_heads, head_size]
175 q_shape_value = self.model.get_constant_value(reshape_q.input[1])
176 if q_shape_value is None:
177 concat = self.model.get_parent(reshape_q, 1)
178 if concat is not None and concat.op_type == "Concat":
179 return self.get_num_heads_and_hidden_size_from_concat(concat)
180 logger.debug("%s is not initializer.", reshape_q.input[1])
181 return self.num_heads, self.hidden_size # Fall back to user specified value
182
183 if (
184 (not isinstance(q_shape_value, np.ndarray))
185 or len(q_shape_value) != 4
186 or (q_shape_value[2] <= 0 or q_shape_value[3] <= 0)
187 ):
188 logger.debug("q_shape_value=%s. Expected value are like [0, 0, num_heads, head_size].", q_shape_value)
189 return self.num_heads, self.hidden_size # Fall back to user specified value
190
191 num_heads = q_shape_value[2]
192 head_size = q_shape_value[3]
193 hidden_size = num_heads * head_size
194
195 if self.num_heads > 0 and num_heads != self.num_heads:
196 if self.num_heads_warning:
197 logger.warning(
198 "--num_heads is %d. Detected value is %d. Using detected value.", self.num_heads, num_heads
199 )
200 self.num_heads_warning = False # Do not show the warning more than once
201
202 if self.hidden_size > 0 and hidden_size != self.hidden_size:
203 if self.hidden_size_warning:
204 logger.warning(
205 "--hidden_size is %d. Detected value is %d. Using detected value.", self.hidden_size, hidden_size
206 )
207 self.hidden_size_warning = False # Do not show the warning more than once
208
209 return num_heads, hidden_size
210
211 def get_add_qk_str(self, add_qk: NodeProto):
212 if not self.shape_infer_done:
213 self.shape_infer = self.model.infer_runtime_shape(update=True)
214 self.shape_infer_done = True
215
216 if self.shape_infer is None:
217 return None
218
219 input_0_shape = self.shape_infer.get_edge_shape(add_qk.input[0])
220 input_1_shape = self.shape_infer.get_edge_shape(add_qk.input[1])
221
222 if input_0_shape is None or input_1_shape is None:
223 logger.debug("one of the inputs of %s is None", add_qk)
224 return None
225
226 if input_0_shape != input_1_shape:
227 logger.debug("the shape of two inputs of %s is not same", add_qk)
228 return None
229
230 return add_qk.input[1]
231
232 def reshape_add_qk(self, add_qk: str):
233 # Convert 4D mask from (B,1,S,T) to (B,N,S,T)
234 # B = batch size, N = num heads, S = source sequence length, T = target sequence length
235 mask_output_name = add_qk + "_mask"
236
237 # Check if concat node for (B,1,S,T) --> (B,N,S,T) already exists
238 concat_node = list(filter(lambda node: node.output[0] == mask_output_name, self.nodes_to_add))
239 if len(concat_node) == 1:
240 return mask_output_name
241
242 assert len(concat_node) == 0
243 concat_node_name = self.model.create_node_name("Concat")
244 concat_add_qk_fp32 = helper.make_node(
245 "Concat",
246 inputs=[add_qk for _ in range(self.num_heads)],
247 outputs=[mask_output_name],
248 name=concat_node_name,
249 axis=1,
250 )
251 # Add new node to graph
252 self.nodes_to_add.append(concat_add_qk_fp32)
253 self.node_name_to_graph_name[concat_node_name] = self.this_graph_name
254
255 return mask_output_name
256
257 def concat_kv(self, past_k: str, past_v: str) -> str:
258 """Concatenate past_k and past_v inputs to create past_kv input.
259
260 Args:
261 past_k (str): name of past K value
262 past_v (str): name of past V value
263
264 Returns:
265 kv_output_name (str): name of past KV value
266 """
267 # Unsqueeze K and V nodes from (B,N,P,H) to (1,B,N,P,H)
268 # B = batch size, N = num heads, P = past sequence length, H = head size
269 unsqueeze_k_name = self.model.create_node_name("Unsqueeze")
270 unsqueeze_v_name = self.model.create_node_name("Unsqueeze")
271 k_5d_name = (past_k + "_5d").replace(".", "_")
272 v_5d_name = (past_v + "_5d").replace(".", "_")
273
274 k_5d = helper.make_node(
275 "Unsqueeze",
276 inputs=[past_k],
277 outputs=[k_5d_name],
278 name=unsqueeze_k_name,
279 axes=[0],
280 )
281 v_5d = helper.make_node(
282 "Unsqueeze",
283 inputs=[past_v],
284 outputs=[v_5d_name],
285 name=unsqueeze_v_name,
286 axes=[0],
287 )
288
289 # Add unsqueeze nodes to graph
290 self.nodes_to_add.append(k_5d)
291 self.nodes_to_add.append(v_5d)
292 self.node_name_to_graph_name[unsqueeze_k_name] = self.this_graph_name
293 self.node_name_to_graph_name[unsqueeze_v_name] = self.this_graph_name
294
295 # Concat K and V to get one node of size (2,B,N,P,H)
296 concat_node_name = self.model.create_node_name("Concat")
297 kv_output_name = past_v.replace(".value", ".kv").replace(".", "_").replace("_value", "_kv")
298 concat_kv = helper.make_node(
299 "Concat",
300 inputs=[k_5d_name, v_5d_name],
301 outputs=[kv_output_name],
302 name=concat_node_name,
303 axis=0,
304 )
305
306 # Add concat node to graph
307 self.nodes_to_add.append(concat_kv)
308 self.node_name_to_graph_name[concat_node_name] = self.this_graph_name
309
310 return kv_output_name
311
312 def split_kv(self, present_k_name: str, present_v_name: str, kv_node: str):
313 """Split kv_node containing present KV values into separate present K and present V values.
314
315 Args:
316 present_k_name (str): name of output to store present K value in
317 present_v_name (str): name of output to store present V value in
318 kv_node (str): name of present KV values
319 """
320 # Split kv_node into present_k and present_v nodes
321
322 # Create initializers for indexing kv_node, whose shape is (2,B,N,P,H)
323 k_index, v_index = "index_0", "index_1"
324 k_dim = self.model.get_initializer(k_index)
325 v_dim = self.model.get_initializer(v_index)
326 if k_dim is None:
327 k_dim = numpy_helper.from_array(np.array(0, dtype="int64"), name=k_index)
328 self.model.add_initializer(k_dim, self.this_graph_name)
329 if v_dim is None:
330 v_dim = numpy_helper.from_array(np.array(1, dtype="int64"), name=v_index)
331 self.model.add_initializer(v_dim, self.this_graph_name)
332
333 # Create nodes to index kv_node
334 gather_k_name = self.model.create_node_name("Gather")
335 gather_v_name = self.model.create_node_name("Gather")
336 present_k = helper.make_node(
337 "Gather",
338 inputs=[kv_node, k_index],
339 outputs=[present_k_name],
340 name=gather_k_name,
341 axis=0,
342 )
343 present_v = helper.make_node(
344 "Gather",
345 inputs=[kv_node, v_index],
346 outputs=[present_v_name],
347 name=gather_v_name,
348 axis=0,
349 )
350
351 # Add gather nodes to graph
352 self.nodes_to_add.append(present_k)
353 self.nodes_to_add.append(present_v)
354 self.node_name_to_graph_name[gather_k_name] = self.this_graph_name
355 self.node_name_to_graph_name[gather_v_name] = self.this_graph_name
356
357 def create_combined_qkv_bias(
358 self,
359 q_add: NodeProto,
360 k_add: NodeProto | None,
361 v_add: NodeProto | None,
362 name_prefix: str,
363 ) -> NodeProto | None:
364 q_bias = self.model.get_initializer(q_add.input[1]) or self.model.get_initializer(q_add.input[0])
365 qb = NumpyHelper.to_array(q_bias)
366 kb = np.zeros_like(qb)
367 vb = np.zeros_like(qb)
368 if k_add is not None:
369 k_bias = self.model.get_initializer(k_add.input[1]) or self.model.get_initializer(k_add.input[0])
370 kb = NumpyHelper.to_array(k_bias)
371 if v_add is not None:
372 v_bias = self.model.get_initializer(v_add.input[1]) or self.model.get_initializer(v_add.input[0])
373 vb = NumpyHelper.to_array(v_bias)
374
375 qkv_bias = np.stack((qb, kb, vb), axis=0)
376 qkv_bias_dim = 3 * np.prod(qb.shape)
377
378 bias_name = name_prefix + "_qkv_bias"
379 self.add_initializer(
380 name=bias_name,
381 data_type=q_bias.data_type,
382 dims=[qkv_bias_dim],
383 vals=qkv_bias,
384 )
385 return bias_name
386
387 def create_packed_qkv_matmul_node(
388 self,
389 q_matmul: NodeProto,
390 k_matmul: NodeProto,
391 v_matmul: NodeProto,
392 q_add: NodeProto,
393 k_add: NodeProto | None,
394 v_add: NodeProto | None,
395 ) -> tuple[NodeProto, NodeProto, NodeProto]:
396 """Create packed QKV MatMul node before MultiHeadAttention node.
397 This is for the scenario where an Attention node should be created but cannot be created
398 because past_key and past_value are separate inputs and not one concatenated input.
399
400 Args:
401 q_matmul (NodeProto): name of MatMul from Q path - (batch_size, sequence_length, hidden_size)
402 k_matmul (NodeProto): name of MatMul from K path - (batch_size, sequence_length, hidden_size)
403 v_matmul (NodeProto): name of MatMul from V path - (batch_size, sequence_length, hidden_size)
404 q_add (NodeProto): name of Add from Q path
405 k_add (NodeProto): name of Add from K path
406 v_add (NodeProto): name of Add from V path
407
408 Returns:
409 q_output (NodeProto): Slice node for Q
410 k_output (NodeProto): Slice node for K
411 v_output (NodeProto): Slice node for V
412 """
413 matmul_node_name = self.model.create_node_name("MatMul")
414
415 # Check that input for Q, K, V is the same
416 assert q_matmul.input[0] == k_matmul.input[0] and k_matmul.input[0] == v_matmul.input[0]
417
418 # Created packed QKV weight
419 q_weight = self.model.get_initializer(q_matmul.input[1])
420 k_weight = self.model.get_initializer(k_matmul.input[1])
421 v_weight = self.model.get_initializer(v_matmul.input[1])
422
423 qw = NumpyHelper.to_array(q_weight)
424 kw = NumpyHelper.to_array(k_weight)
425 vw = NumpyHelper.to_array(v_weight)
426
427 assert qw.shape == kw.shape and kw.shape == vw.shape
428 d = qw.shape[0]
429
430 qkv_weight = np.stack((qw, kw, vw), axis=1).reshape((d, 3 * d))
431 qkv_weight_name = matmul_node_name + "_qkv_weight"
432
433 self.add_initializer(
434 name=qkv_weight_name,
435 data_type=q_weight.data_type,
436 dims=[qkv_weight.shape[0], qkv_weight.shape[1]],
437 vals=qkv_weight,
438 )
439
440 # Created packed QKV MatMul with output (B, S, 3*D)
441 # Output is of the form:
442 #
443 # [[[Q Q ... Q Q K K ... K K V V ... V V]]]
444 # [Q Q ... Q Q K K ... K K V V ... V V]
445 # .
446 # .
447 # .
448 # [[Q Q ... Q Q K K ... K K V V ... V V]
449 # [Q Q ... Q Q K K ... K K V V ... V V]]]
450 qkv_matmul_output = matmul_node_name + "_qkv_out"
451 qkv_matmul = helper.make_node(
452 "MatMul",
453 inputs=[q_matmul.input[0], qkv_weight_name],
454 outputs=[qkv_matmul_output],
455 name=matmul_node_name,
456 )
457 self.node_name_to_graph_name[matmul_node_name] = self.this_graph_name
458
459 qkv_nodes = [qkv_matmul]
460
461 # Create Slice nodes to access Q, K, V
462 q_slice_name = matmul_node_name + "_q_start_index"
463 self.add_initializer(name=q_slice_name, data_type=TensorProto.INT64, dims=[1], vals=[0], raw=False)
464 k_slice_name = matmul_node_name + "_k_start_index"
465 self.add_initializer(name=k_slice_name, data_type=TensorProto.INT64, dims=[1], vals=[d], raw=False)
466 v_slice_name = matmul_node_name + "_v_start_index"
467 self.add_initializer(name=v_slice_name, data_type=TensorProto.INT64, dims=[1], vals=[2 * d], raw=False)
468 end_of_qkv_name = matmul_node_name + "_end_of_qkv_index"
469 self.add_initializer(name=end_of_qkv_name, data_type=TensorProto.INT64, dims=[1], vals=[3 * d], raw=False)
470 qkv_last_axis_name = matmul_node_name + "_qkv_last_axis"
471 self.add_initializer(name=qkv_last_axis_name, data_type=TensorProto.INT64, dims=[1], vals=[-1], raw=False)
472
473 q_slice_output = matmul_node_name + "_q_out"
474 q_slice = helper.make_node(
475 "Slice",
476 inputs=[qkv_matmul_output, q_slice_name, k_slice_name, qkv_last_axis_name],
477 outputs=[q_slice_output],
478 name=self.model.create_node_name("Slice"),
479 )
480 self.node_name_to_graph_name[q_slice.name] = self.this_graph_name
481 k_slice_output = matmul_node_name + "_k_out"
482 k_slice = helper.make_node(
483 "Slice",
484 inputs=[qkv_matmul_output, k_slice_name, v_slice_name, qkv_last_axis_name],
485 outputs=[k_slice_output],
486 name=self.model.create_node_name("Slice"),
487 )
488 self.node_name_to_graph_name[k_slice.name] = self.this_graph_name
489 v_slice_output = matmul_node_name + "_v_out"
490 v_slice = helper.make_node(
491 "Slice",
492 inputs=[qkv_matmul_output, v_slice_name, end_of_qkv_name, qkv_last_axis_name],
493 outputs=[v_slice_output],
494 name=self.model.create_node_name("Slice"),
495 )
496 self.node_name_to_graph_name[v_slice.name] = self.this_graph_name
497
498 q_output = q_slice
499 k_output = k_slice
500 v_output = v_slice
501 qkv_nodes.extend([q_slice, k_slice, v_slice])
502
503 if self.disable_multi_head_attention_bias:
504 if q_add is not None:
505 initializer_input = 1 if self.model.get_initializer(q_add.input[1]) else 0
506 if np.any(NumpyHelper.to_array(self.model.get_initializer(q_add.input[initializer_input]))):
507 q_add.input[1 - initializer_input] = q_slice_output
508 q_output = q_add
509 qkv_nodes.append(q_add)
510 self.node_name_to_graph_name[q_add.name] = self.this_graph_name
511 if k_add is not None:
512 initializer_input = 1 if self.model.get_initializer(k_add.input[1]) else 0
513 if np.any(NumpyHelper.to_array(self.model.get_initializer(k_add.input[initializer_input]))):
514 k_add.input[1 - initializer_input] = k_slice_output
515 k_output = k_add
516 qkv_nodes.append(k_add)
517 self.node_name_to_graph_name[k_add.name] = self.this_graph_name
518 if v_add is not None:
519 initializer_input = 1 if self.model.get_initializer(v_add.input[1]) else 0
520 if np.any(NumpyHelper.to_array(self.model.get_initializer(v_add.input[initializer_input]))):
521 v_add.input[1 - initializer_input] = v_slice_output
522 v_output = v_add
523 qkv_nodes.append(v_add)
524 self.node_name_to_graph_name[v_add.name] = self.this_graph_name
525
526 # Add nodes to graph
527 self.nodes_to_add.extend(qkv_nodes)
528 return q_output, k_output, v_output
529
530 # This function is used in child classes for bart or conformer model.
531 def create_multihead_attention_node(
532 self,
533 q_matmul: NodeProto,
534 k_matmul: NodeProto | str | None,
535 v_matmul: NodeProto | str | None,
536 q_add: NodeProto,
537 k_add: NodeProto | None,
538 v_add: NodeProto | None,
539 num_heads: int,
540 hidden_size: int,
541 output: str,
542 key_padding_mask: str = "",
543 add_qk: str = "",
544 unidirectional: bool = False,
545 past_k: str = "",
546 past_v: str = "",
547 present_k: str = "",
548 present_v: str = "",
549 packed_qkv: bool = False,
550 ) -> NodeProto | None:
551 """Create a MultiHeadAttention node.
552
553 Args:
554 q_matmul (NodeProto): name of MatMul from Q path - (batch_size, sequence_length, hidden_size)
555 k_matmul (NodeProto): name of MatMul from K path - (batch_size, sequence_length, hidden_size) or (batch_size, num_heads, past_sequence_length, head_size)
556 v_matmul (NodeProto): name of MatMul from V path - (batch_size, sequence_length, hidden_size) or (batch_size, num_heads, past_sequence_length, head_size)
557 q_add (NodeProto): name of Add from Q path
558 k_add (NodeProto): name of Add from K path
559 v_add (NodeProto): name of Add from V path
560 num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
561 hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
562 output (str): output name of MHA
563 key_padding_mask (str): name of key padding mask
564 add_qk (str): name of add after Q x K'
565 unidirectional (bool): whether to apply causal attention mask automatically or not
566 past_k (str): name of past K value - (batch_size, num_heads, past_sequence_length, head_size)
567 past_v (str): name of past V value - (batch_size, num_heads, past_sequence_length, head_size)
568 present_k (str): name of present K value - (batch_size, num_heads, sequence_length, head_size)
569 present_v (str): name of present V value - (batch_size, num_heads, sequence_length, head_size)
570 packed_qkv (bool): whether to combine MatMuls from Q, K, V paths
571 Note: This is for the scenario where an Attention node should be created but cannot be created
572 because past_key and past_value are separate inputs and not one concatenated input.
573
574 Returns:
575 Union[NodeProto, None]: the node created or None if failed.
576 """
577 # B = batch size, N = num heads, P = past seq len, H = head size
578 assert num_heads > 0
579
580 if hidden_size > 0 and (hidden_size % num_heads) != 0:
581 logger.debug("input hidden size %d is not a multiple of num of heads %d", hidden_size, num_heads)
582 return None
583
584 graph_input_names = {node.name for node in self.model.graph().input}
585 mha_node_name = self.model.create_node_name("Attention")
586
587 # Add initial Q/K/V inputs for MHA
588 mha_inputs = []
589 if packed_qkv:
590 q_slice, k_slice, v_slice = self.create_packed_qkv_matmul_node(
591 q_matmul,
592 k_matmul,
593 v_matmul,
594 q_add,
595 k_add,
596 v_add,
597 )
598 mha_inputs.extend([q_slice.output[0], k_slice.output[0], v_slice.output[0]])
599 elif isinstance(k_matmul, NodeProto) and isinstance(v_matmul, NodeProto):
600 if self.disable_multi_head_attention_bias:
601 mha_inputs.extend([q_add.output[0], k_matmul.output[0], v_add.output[0]])
602 else:
603 mha_inputs.extend([q_matmul.output[0], k_matmul.output[0], v_matmul.output[0]])
604 elif (
605 isinstance(k_matmul, str)
606 and isinstance(v_matmul, str)
607 and k_matmul in graph_input_names
608 and v_matmul in graph_input_names
609 ):
610 if self.disable_multi_head_attention_bias:
611 mha_inputs.extend([q_add.output[0], k_matmul, v_matmul])
612 else:
613 mha_inputs.extend([q_matmul.output[0], k_matmul, v_matmul])
614 else:
615 return None
616
617 # Add bias to inputs for MHA
618 # Bias for cross attention is not fully supported in DMMHA and cpu MHA kernels since they assume
619 # bias has been added to key and value when they are in BNSH format, so only bias for query is used.
620 # Need add checks if we found such assumption is not true.
621 if not self.disable_multi_head_attention_bias:
622 bias_name = self.create_combined_qkv_bias(q_add, k_add, v_add, mha_node_name)
623 mha_inputs.append(bias_name)
624 else:
625 mha_inputs.append("")
626
627 # Add optional inputs for MHA
628 if past_k and past_v:
629 mha_inputs.extend([key_padding_mask, add_qk, past_k, past_v])
630 elif key_padding_mask or add_qk:
631 mha_inputs.extend([key_padding_mask, add_qk])
632
633 # Add outputs for MHA
634 mha_outputs = [output]
635 if present_k and present_v:
636 mha_outputs.extend([present_k, present_v])
637
638 mha_node = helper.make_node(
639 "MultiHeadAttention",
640 inputs=mha_inputs,
641 outputs=mha_outputs,
642 name=mha_node_name,
643 )
644 mha_node.domain = "com.microsoft"
645 mha_node.attribute.append(helper.make_attribute("num_heads", num_heads))
646 if unidirectional:
647 mha_node.attribute.append(helper.make_attribute("unidirectional", int(unidirectional)))
648
649 self.increase_counter("MultiHeadAttention")
650 return mha_node
651
652 def create_attention_node(
653 self,
654 mask_index: str | None,
655 q_matmul: NodeProto,
656 k_matmul: NodeProto,
657 v_matmul: NodeProto,
658 q_add: NodeProto,
659 k_add: NodeProto,
660 v_add: NodeProto,
661 num_heads: int,
662 hidden_size: int,
663 first_input: str,
664 output: str,
665 add_qk_str: str = "",
666 causal: bool = False,
667 past_k: str = "",
668 past_v: str = "",
669 present_k: str = "",
670 present_v: str = "",
671 scale: float | None = None,
672 ) -> NodeProto | None:
673 """Create an Attention node.
674
675 Args:
676 mask_index (str | None): mask input
677 q_matmul (NodeProto): MatMul node in fully connection for Q
678 k_matmul (NodeProto): MatMul node in fully connection for K
679 v_matmul (NodeProto): MatMul node in fully connection for V
680 q_add (NodeProto): Add bias node in fully connection for Q
681 k_add (NodeProto): Add bias node in fully connection for K
682 v_add (NodeProto): Add bias node in fully connection for V
683 num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
684 hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
685 first_input (str): first input name
686 output (str): output name
687 add_qk_str (str): name of Add node after Q x K'
688 causal: whether it is uni-directional mask.
689 past_k (str): name of input for past K value
690 past_v (str): name of input for past V value
691 present_k (str): name of output to store present K value
692 present_v (str): name of output to store present V value
693 scale: scale before softmax
694
695 Returns:
696 Union[NodeProto, None]: the node created or None if failed.
697 """
698 assert num_heads > 0
699
700 if hidden_size > 0 and (hidden_size % num_heads) != 0:
701 logger.debug("input hidden size %d is not a multiple of num of heads %d", hidden_size, num_heads)
702 return None
703
704 has_bias = True
705 if q_add is None and k_add is None and v_add is None:
706 has_bias = False
707
708 q_weight = self.model.get_initializer(q_matmul.input[1])
709 k_weight = self.model.get_initializer(k_matmul.input[1])
710 v_weight = self.model.get_initializer(v_matmul.input[1])
711
712 q_bias, k_bias, v_bias = None, None, None
713 if has_bias:
714 q_bias = self.model.get_initializer(q_add.input[1]) or self.model.get_initializer(q_add.input[0])
715 k_bias = self.model.get_initializer(k_add.input[1]) or self.model.get_initializer(k_add.input[0])
716 v_bias = self.model.get_initializer(v_add.input[1]) or self.model.get_initializer(v_add.input[0])
717
718 if not (k_weight and v_weight and q_bias and k_bias):
719 return None
720
721 if q_weight is None:
722 print(
723 f"{q_matmul.input[1]} is not an initializer. "
724 "Please set do_constant_folding=True in torch.onnx.export to unblock attention fusion"
725 )
726 return None
727
728 qw = NumpyHelper.to_array(q_weight)
729 kw = NumpyHelper.to_array(k_weight)
730 vw = NumpyHelper.to_array(v_weight)
731
732 # assert q and k have same shape as expected
733 assert qw.shape == kw.shape
734
735 qw_in_size = qw.shape[0]
736 kw_in_size = kw.shape[0]
737 vw_in_size = vw.shape[0]
738
739 assert qw_in_size == kw_in_size == vw_in_size
740
741 if hidden_size > 0 and hidden_size != qw_in_size:
742 logger.warning(
743 "Input hidden size (%d) is not same as weight matrix dimension of q,k,v (%d). "
744 "Please provide a correct input hidden size or pass in 0",
745 hidden_size,
746 qw_in_size,
747 )
748
749 is_qkv_diff_dims = False
750 if qw.shape != vw.shape:
751 is_qkv_diff_dims = True
752
753 # All the matrices can have the same shape or q, k matrices can have the same shape with v being different
754 # For 2d weights, the shapes would be [in_size, out_size].
755 # For 3d weights, shape would be [in_size, a, b] where a*b = out_size
756 qw_out_size = np.prod(qw.shape[1:])
757 kw_out_size = np.prod(kw.shape[1:])
758 vw_out_size = np.prod(vw.shape[1:])
759
760 qkv_weight_dim = 0
761 if is_qkv_diff_dims:
762 qkv_weight = np.concatenate((qw, kw, vw), axis=1)
763 qkv_weight_dim = qw_out_size + kw_out_size + vw_out_size
764 else:
765 qkv_weight = np.stack((qw, kw, vw), axis=1)
766 qkv_weight_dim = 3 * qw_out_size
767
768 qkv_bias_dim = 0
769 qkv_bias: np.ndarray | None = None
770 if has_bias:
771 qb = NumpyHelper.to_array(q_bias)
772 kb = NumpyHelper.to_array(k_bias)
773 vb = NumpyHelper.to_array(v_bias)
774
775 q_bias_shape = np.prod(qb.shape)
776 k_bias_shape = np.prod(kb.shape)
777 v_bias_shape = np.prod(vb.shape)
778
779 assert q_bias_shape == k_bias_shape == qw_out_size
780 assert v_bias_shape == vw_out_size
781
782 if is_qkv_diff_dims:
783 qkv_bias = np.concatenate((qb, kb, vb), axis=0)
784 qkv_bias_dim = q_bias_shape + k_bias_shape + v_bias_shape
785 else:
786 qkv_bias = np.stack((qb, kb, vb), axis=0)
787 qkv_bias_dim = 3 * q_bias_shape
788
789 attention_node_name = self.model.create_node_name("Attention")
790
791 if not self.use_multi_head_attention:
792 self.add_initializer(
793 name=attention_node_name + "_qkv_weight",
794 data_type=q_weight.data_type,
795 dims=[qw_in_size, int(qkv_weight_dim)],
796 vals=qkv_weight,
797 )
798
799 if has_bias:
800 self.add_initializer(
801 name=attention_node_name + "_qkv_bias",
802 data_type=q_bias.data_type,
803 dims=[int(qkv_bias_dim)],
804 vals=qkv_bias,
805 )
806
807 # For MultiHeadAttention operator, use separated inputs for query, key and value, and no weights.
808 if self.use_multi_head_attention:
809 if add_qk_str:
810 logger.debug("MultiHeadAttention does not support relative_position_bias: cannot fuse the attention.")
811 return None
812
813 attention_inputs = [
814 q_matmul.output[0],
815 k_matmul.output[0],
816 v_matmul.output[0],
817 attention_node_name + "_qkv_bias",
818 ]
819
820 if mask_index is not None:
821 attention_inputs.append(mask_index)
822
823 attention_node = helper.make_node(
824 "MultiHeadAttention",
825 inputs=attention_inputs,
826 outputs=[output],
827 name=attention_node_name,
828 )
829 self.increase_counter("MultiHeadAttention")
830
831 else:
832 attention_inputs = [
833 first_input,
834 attention_node_name + "_qkv_weight",
835 attention_node_name + "_qkv_bias" if has_bias else "",
836 ]
837 if mask_index is not None:
838 attention_inputs.append(mask_index)
839 else:
840 attention_inputs.append("")
841
842 past_exists = past_k and past_v
843 if past_exists:
844 past_kv = self.concat_kv(past_k, past_v)
845 attention_inputs.append(past_kv)
846
847 if add_qk_str:
848 # Add additional add to attention node (input name = attention_bias)
849 if not past_exists:
850 attention_inputs.append("")
851 attention_inputs.append(add_qk_str)
852
853 attention_outputs = [output]
854 if present_k and present_v:
855 present_kv = present_k.replace(".key", "").replace("_key", "").replace(".", "_")
856 attention_outputs.append(present_kv)
857 self.split_kv(present_k, present_v, present_kv)
858
859 attention_node = helper.make_node(
860 "Attention",
861 inputs=attention_inputs,
862 outputs=attention_outputs,
863 name=attention_node_name,
864 )
865 self.increase_counter("Attention")
866
867 attention_node.domain = "com.microsoft"
868 attention_node.attribute.extend([helper.make_attribute("num_heads", num_heads)])
869
870 if causal:
871 attention_node.attribute.extend([helper.make_attribute("unidirectional", 1)])
872
873 if scale is not None:
874 attention_node.attribute.extend([helper.make_attribute("scale", scale)])
875
876 if is_qkv_diff_dims:
877 attention_node.attribute.extend(
878 [helper.make_attribute("qkv_hidden_sizes", [qw_out_size, kw_out_size, vw_out_size])]
879 )
880
881 if self.mask_filter_value is not None:
882 attention_node.attribute.extend([helper.make_attribute("mask_filter_value", float(self.mask_filter_value))])
883
884 return attention_node
885
886 def fuse(self, node, input_name_to_nodes, output_name_to_node):
887 # Sometimes we can not fuse skiplayernormalization since the add before layernorm has an output that used by nodes outside skiplayernorm
888 # Conceptually we treat add before layernorm as skiplayernorm node since they share the same pattern
889 normalize_node = node
890 start_node = normalize_node
891 if normalize_node.op_type == "LayerNormalization":
892 add_before_layernorm = self.model.match_parent(normalize_node, "Add", 0)
893 if add_before_layernorm is not None:
894 start_node = add_before_layernorm
895 elif self.model.find_graph_input(normalize_node.input[0]) is not None:
896 # Pre-LN first block: LN fed directly by graph input. QKV matching will
897 # still fail from this (first) LN anchor because its inputs are weights, not
898 # the QKV projection path. The real fusion happens when fuse() is called
899 # again from the second LN/SkipLN anchor after the residual Add, where the
900 # other_inputs and root_input changes (#2-#4) take effect.
901 start_node = normalize_node
902 else:
903 return
904
905 # SkipLayerNormalization has two inputs, and one of them is the root input for attention.
906 qkv_nodes = self.model.match_parent_path(
907 start_node,
908 ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
909 [None, None, 0, 0, 0],
910 )
911 einsum_node = None
912 if qkv_nodes is not None:
913 (_, _, reshape_qkv, transpose_qkv, matmul_qkv) = qkv_nodes
914 else:
915 # Match Albert
916 qkv_nodes = self.model.match_parent_path(
917 start_node, ["Add", "Einsum", "Transpose", "MatMul"], [1, None, 0, 0]
918 )
919 if qkv_nodes is not None:
920 (_, einsum_node, transpose_qkv, matmul_qkv) = qkv_nodes
921 else:
922 return
923
924 other_inputs = []
925 for _i, node_input in enumerate(start_node.input):
926 if node_input not in output_name_to_node:
927 if self.model.find_graph_input(node_input) is None:
928 continue
929
930 if node_input == qkv_nodes[0].output[0]:
931 continue
932 other_inputs.append(node_input)
933 if len(other_inputs) != 1:
934 return
935
936 root_input = other_inputs[0]
937
938 # Match flaubert Mask
939 # |
940 # Mul --> LayerNormalization --> Attention --> MatMul --> Add
941 # | |
942 # | |
943 # +---------------------------------------------------------
944 mul_before_layernorm = self.model.match_parent(start_node, "Mul", 0)
945 if mul_before_layernorm is not None:
946 mul_children = input_name_to_nodes[mul_before_layernorm.output[0]]
947 if mul_children is not None and len(mul_children) == 2:
948 layernorm_node = mul_children[1]
949 if layernorm_node.op_type == "LayerNormalization":
950 root_input = layernorm_node.output[0]
951 else:
952 return
953 elif mul_children is not None and len(mul_children) == 5:
954 root_input = mul_before_layernorm.output[0]
955 else:
956 return
957 elif normalize_node.op_type in ("LayerNormalization", "SkipLayerNormalization"):
958 children = input_name_to_nodes[root_input]
959 for child in children:
960 if child.op_type == "LayerNormalization":
961 root_input = child.output[0]
962
963 # When Add before the LayerNormalization produces an output
964 # that is consumed by some other nodes other than the LayerNormalization itself,
965 # fused SkipLayerNormalization will have several outputs.
966 # In this case we need to pick the one used in Attention
967 # For example, this is the case for ViT
968 # SkipLayerNormalization --> Attention --> MatMul --> Add --> SkipLayerNormalization
969 # | |
970 # | |
971 # +---------------------------------------------------------------------+
972 if root_input in output_name_to_node:
973 parent_node = output_name_to_node[root_input]
974 if parent_node.op_type == "SkipLayerNormalization" and len(parent_node.output) == 4:
975 root_input = parent_node.output[0]
976
977 children = input_name_to_nodes[root_input]
978 children_types = [child.op_type for child in children]
979 if children_types.count("MatMul") != 3:
980 return
981
982 v_nodes = self.model.match_parent_path(matmul_qkv, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, None])
983 if v_nodes is None:
984 logger.debug("fuse_attention: failed to match v path")
985 return
986 (_, _, add_v, matmul_v) = v_nodes
987
988 is_distill = False
989 is_distill_add = False
990 is_no_mask_attention = False
991 is_sdpa = False
992 qk_paths = {
993 "path1": (["Softmax", "Add", "Div", "MatMul"], [0, 0, None, 0]),
994 "path2": (["Softmax", "Add", "Mul", "MatMul"], [0, 0, None, 0]),
995 "path3": (["Softmax", "Where", "MatMul", "Div"], [0, 0, 2, 0]),
996 "path4": (["Softmax", "Add", "Where", "MatMul"], [0, 0, 0, 2]),
997 "path5": (["Softmax", "Div", "MatMul"], [0, 0, 0]),
998 "sdpa": (["Softmax", "Add", "MatMul", "Mul", "Sqrt"], [0, 0, None, 0, 1]),
999 }
1000
1001 qk_nodes = None
1002 for k, v in qk_paths.items():
1003 qk_nodes = self.model.match_parent_path(matmul_qkv, v[0], v[1])
1004 if qk_nodes is None:
1005 continue
1006 if k == "path3":
1007 is_distill = True
1008 elif k == "path4":
1009 is_distill_add = True
1010 elif k == "path5":
1011 is_no_mask_attention = True
1012 elif k == "sdpa":
1013 is_sdpa = True
1014 break
1015
1016 if qk_nodes is None:
1017 logger.debug("fuse_attention: failed to match qk path")
1018 return
1019
1020 add_qk = None
1021 matmul_qk = None
1022 where_qk = None
1023 after_q = None
1024 if is_distill:
1025 (_, where_qk, matmul_qk, _) = qk_nodes
1026 elif is_distill_add:
1027 (_, add_qk, where_qk, matmul_qk) = qk_nodes
1028 elif is_no_mask_attention:
1029 (_, _, matmul_qk) = qk_nodes
1030 elif is_sdpa:
1031 (_, add_qk, matmul_qk, after_q, _) = qk_nodes
1032 else:
1033 (_, add_qk, _, matmul_qk) = qk_nodes
1034
1035 after_q = after_q or matmul_qk
1036 q_nodes = self.model.match_parent_path(after_q, ["Transpose", "Reshape", "Add", "MatMul"], [0, 0, 0, None])
1037 if q_nodes is None:
1038 q_nodes = self.model.match_parent_path(
1039 after_q,
1040 ["Div", "Transpose", "Reshape", "Add", "MatMul"],
1041 [0, 0, 0, 0, None],
1042 )
1043 if q_nodes is None:
1044 logger.debug("fuse_attention: failed to match q path")
1045 return
1046 reshape_q = q_nodes[-3]
1047 add_q = q_nodes[-2]
1048 matmul_q = q_nodes[-1]
1049
1050 after_k = matmul_qk
1051 if is_sdpa:
1052 mul_k_nodes = self.model.match_parent_path(matmul_qk, ["Mul", "Sqrt"], [1, None])
1053 if mul_k_nodes is None:
1054 logger.debug("fuse_attention: failed to match mul sqrt q path")
1055 return
1056 (after_k, _) = mul_k_nodes
1057
1058 k_nodes = self.model.match_parent_path(
1059 after_k, ["Transpose", "Reshape", "Add", "MatMul"], [0 if is_sdpa else 1, 0, 0, None]
1060 )
1061 if k_nodes is None:
1062 k_nodes = self.model.match_parent_path(
1063 matmul_qk,
1064 ["Transpose", "Transpose", "Reshape", "Add", "MatMul"],
1065 [1, 0, 0, 0, None],
1066 )
1067 if k_nodes is None:
1068 logger.debug("fuse_attention: failed to match k path")
1069 return
1070 add_k = k_nodes[-2]
1071 matmul_k = k_nodes[-1]
1072
1073 # Note that Cast might be removed by OnnxRuntime so we match two patterns here.
1074 mask_nodes = None
1075 add_qk_str = ""
1076 if is_distill:
1077 _, mask_nodes, _ = self.model.match_parent_paths(
1078 where_qk,
1079 [
1080 (["Expand", "Reshape", "Equal"], [0, 0, 0]),
1081 (["Equal", "Unsqueeze", "Unsqueeze"], [0, 0, 0]),
1082 (["Cast", "Expand", "Reshape", "Equal"], [0, 0, 0, 0]),
1083 ],
1084 output_name_to_node,
1085 )
1086 elif is_distill_add:
1087 _, mask_nodes, _ = self.model.match_parent_paths(
1088 where_qk,
1089 [
1090 (["Cast", "Equal", "Unsqueeze", "Unsqueeze"], [0, 0, 0, 0]),
1091 (["Equal", "Unsqueeze", "Unsqueeze"], [0, 0, 0]),
1092 ],
1093 output_name_to_node,
1094 )
1095 if add_qk is not None:
1096 add_qk_str = self.get_add_qk_str(add_qk)
1097 if add_qk_str is None:
1098 logger.debug("fuse_attention: failed to verify shape inference of %s", add_qk)
1099 return
1100 elif is_no_mask_attention:
1101 pass
1102 else:
1103 _, mask_nodes, _ = self.model.match_parent_paths(
1104 add_qk,
1105 [
1106 (["Mul", "Sub", "Cast", "Unsqueeze", "Unsqueeze"], [None, 0, 1, 0, 0]),
1107 (["Mul", "Sub", "Unsqueeze", "Unsqueeze"], [None, 0, 1, 0]),
1108 # The following two patterns are for SDPA.
1109 (["Where", "Cast", "Sub", "Expand", "Unsqueeze", "Unsqueeze"], [None, 0, 0, 1, 0, 0]),
1110 (["Where", "Cast", "Sub", "Cast", "Expand", "Unsqueeze", "Unsqueeze"], [None, 0, 0, 1, 0, 0, 0]),
1111 ],
1112 output_name_to_node,
1113 )
1114 if not is_no_mask_attention and mask_nodes is None:
1115 logger.debug("fuse_attention: failed to match mask path")
1116 return
1117
1118 if not is_no_mask_attention and len(mask_nodes) > 1:
1119 _, mul_val = self.model.get_constant_input(mask_nodes[0])
1120 # The mask value shall be a float scalar (usually is the lowest float value).
1121 if (
1122 (mul_val is None)
1123 or not (isinstance(mul_val, np.ndarray) and mul_val.size == 1)
1124 or (mul_val.item() >= 0)
1125 ):
1126 return
1127 if mul_val.item() != -10000:
1128 self.mask_filter_value = mul_val.item()
1129
1130 if matmul_v.input[0] == root_input and matmul_q.input[0] == root_input and matmul_k.input[0] == root_input:
1131 mask_index = self.attention_mask.process_mask(mask_nodes[-1].input[0]) if not is_no_mask_attention else None
1132
1133 attention_last_node = reshape_qkv if einsum_node is None else transpose_qkv
1134
1135 q_num_heads, q_hidden_size = self.get_num_heads_and_hidden_size(reshape_q)
1136 if q_num_heads <= 0 or q_hidden_size <= 0:
1137 logger.warning(
1138 "Failed to detect num_heads and hidden_size for Attention fusion. "
1139 "Please specify those parameters in argument."
1140 )
1141 return
1142
1143 # number of heads are same for all the paths, hence to create attention node, we pass the q_num_heads
1144 # the input_hidden_size represents the input hidden size, this is used as needed but hidden sizes for Q, K are extracted appropriately
1145 new_node = self.create_attention_node(
1146 mask_index=mask_index,
1147 q_matmul=matmul_q,
1148 k_matmul=matmul_k,
1149 v_matmul=matmul_v,
1150 q_add=add_q,
1151 k_add=add_k,
1152 v_add=add_v,
1153 num_heads=q_num_heads,
1154 hidden_size=q_hidden_size,
1155 first_input=root_input,
1156 output=attention_last_node.output[0],
1157 add_qk_str=add_qk_str,
1158 )
1159
1160 if new_node is None:
1161 return
1162
1163 self.nodes_to_add.append(new_node)
1164 self.node_name_to_graph_name[new_node.name] = self.this_graph_name
1165
1166 if einsum_node is not None:
1167 unique_index = einsum_node.input[0]
1168 new_edge = "edge_modified_" + unique_index
1169
1170 shape_tensor = self.add_initializer(
1171 name="shape_modified_tensor" + unique_index,
1172 data_type=TensorProto.INT64,
1173 dims=[4],
1174 vals=[0, 0, q_num_heads, int(q_hidden_size / q_num_heads)],
1175 raw=False,
1176 )
1177
1178 self.model.add_node(
1179 helper.make_node(
1180 "Reshape",
1181 [attention_last_node.output[0], shape_tensor.name],
1182 [new_edge],
1183 "reshape_modified_" + unique_index,
1184 ),
1185 self.this_graph_name,
1186 )
1187 einsum_node.input[0] = new_edge
1188
1189 self.nodes_to_remove.extend([attention_last_node, transpose_qkv, matmul_qkv])
1190 self.nodes_to_remove.extend(qk_nodes)
1191
1192 # For MultiHeadAttention operator, MatMul nodes for Q/K/V projection shall not be fused.
1193 self.nodes_to_remove.extend(q_nodes if not self.use_multi_head_attention else q_nodes[:-1])
1194 self.nodes_to_remove.extend(k_nodes if not self.use_multi_head_attention else k_nodes[:-1])
1195 self.nodes_to_remove.extend(v_nodes if not self.use_multi_head_attention else v_nodes[:-1])
1196
1197 # Use prune graph to remove mask nodes since they are shared by all attention nodes.
1198 self.prune_graph = True
1199 