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
7from fusion_attention import AttentionMask, FusionAttention
8from fusion_options import AttentionMaskFormat
9from onnx import NodeProto
10from onnx_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionAttentionClip(FusionAttention):
16 """
17 Fuse Attention subgraph of Clip into one Attention node.
18 """
19
20 def __init__(
21 self,
22 model: OnnxModel,
23 hidden_size: int,
24 num_heads: int,
25 ):
26 attention_mask = AttentionMask(model)
27 attention_mask.mask_format = AttentionMaskFormat.NoMask
28
29 super().__init__(
30 model,
31 hidden_size,
32 num_heads,
33 attention_mask,
34 use_multi_head_attention=False,
35 search_op_types=["SkipLayerNormalization"],
36 )
37
38 def get_num_heads_and_hidden_size(self, reshape_q: NodeProto) -> tuple[int, int]:
39 """Detect num_heads and hidden_size for ONNX model from MiDaS
40 Args:
41 reshape_q (NodeProto): reshape node for q
42 Returns:
43 Tuple[int, int]: num_heads and hidden_size
44 """
45 concat = self.model.match_parent(reshape_q, "Concat", 1)
46 if concat is None or len(concat.input) != 4:
47 return self.num_heads, self.hidden_size
48
49 # The shape is a tensor like [?, ?, num_heads, head_size]
50 num_head_value = self.model.get_constant_value(concat.input[2])
51 if num_head_value is None:
52 return self.num_heads, self.hidden_size # Fall back to user specified value
53
54 if len(num_head_value) != 1 or num_head_value[0] <= 0:
55 return self.num_heads, self.hidden_size # Fall back to user specified value
56
57 num_heads = num_head_value[0]
58
59 head_size_value = self.model.get_constant_value(concat.input[3])
60 if head_size_value is None:
61 return self.num_heads, self.hidden_size # Fall back to user specified value
62
63 if len(head_size_value) != 1 or head_size_value[0] <= 0:
64 return self.num_heads, self.hidden_size # Fall back to user specified value
65
66 head_size = head_size_value[0]
67
68 hidden_size = num_heads * head_size
69
70 if self.num_heads > 0 and num_heads != self.num_heads:
71 if self.num_heads_warning:
72 logger.warning(f"--num_heads is {self.num_heads}. Detected value is {num_heads}. Using detected value.")
73 self.num_heads_warning = False # Do not show the warning more than once
74
75 if self.hidden_size > 0 and hidden_size != self.hidden_size:
76 if self.hidden_size_warning:
77 logger.warning(
78 f"--hidden_size is {self.hidden_size}. Detected value is {hidden_size}. Using detected value."
79 )
80 self.hidden_size_warning = False # Do not show the warning more than once
81
82 return num_heads, hidden_size
83
84 def fuse(self, normalize_node, input_name_to_nodes, output_name_to_node):
85 skip_input_index = None
86 node_before_layer_norm = None
87 for i in [1, 0]:
88 parent = self.model.match_parent(normalize_node, "SkipLayerNormalization", i)
89 if parent is not None:
90 skip_input_index = i
91 node_before_layer_norm = parent
92
93 root_input = None
94 if node_before_layer_norm is not None:
95 root_input = node_before_layer_norm.output[0]
96 else:
97 # Deal with the first attention after the embedding layer.
98 for i in [0, 1]:
99 node_before_layer_norm = None
100
101 node_before_layer_norm_1 = self.model.match_parent(normalize_node, "Add", i)
102 node_before_layer_norm_2 = self.model.match_parent(normalize_node, "LayerNormalization", i)
103 if node_before_layer_norm_1 is not None:
104 # Add -----------+
105 # | |
106 # LayerNorm |
107 # | |
108 # LayerNorm |
109 # | |
110 # Attention subgraph |
111 # | |
112 # SkipLayerNorm ------+
113 node_before_layer_norm = node_before_layer_norm_1
114 elif node_before_layer_norm_2 is not None:
115 # Add
116 # |
117 # LayerNorm --------+
118 # | |
119 # LayerNorm |
120 # | |
121 # Attention subgraph |
122 # | |
123 # SkipLayerNorm ------+
124 node_before_layer_norm = node_before_layer_norm_2
125
126 if node_before_layer_norm is None:
127 continue
128 child = self.model.find_first_child_by_type(
129 node_before_layer_norm,
130 "LayerNormalization",
131 input_name_to_nodes,
132 False,
133 )
134 if child is None:
135 continue
136 root_input = child.output[0]
137 skip_input_index = i
138 break
139
140 if skip_input_index is None:
141 return
142
143 qkv_nodes = self.model.match_parent_path(
144 normalize_node,
145 ["Add", "MatMul", "Reshape", "Transpose", "Reshape", "MatMul"],
146 [1 - skip_input_index, None, None, 0, 0, 0],
147 )
148 if qkv_nodes is None:
149 qkv_nodes = self.model.match_parent_path(
150 normalize_node,
151 ["Add", "MatMul", "Reshape", "Transpose", "MatMul"],
152 [1, None, 0, 0, 0],
153 )
154 if qkv_nodes is None:
155 logger.debug("fuse_attention: failed to match qkv path")
156 return
157 reshape_qkv, transpose_qkv, matmul_qkv = (
158 qkv_nodes[2],
159 qkv_nodes[3],
160 qkv_nodes[-1],
161 )
162
163 v_nodes = self.model.match_parent_path(
164 matmul_qkv,
165 ["Reshape", "Transpose", "Reshape", "Add", "MatMul"],
166 [1, 0, 0, 0, None],
167 )
168 if v_nodes is None:
169 v_nodes = self.model.match_parent_path(
170 matmul_qkv, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, None]
171 )
172 if v_nodes is None:
173 logger.debug("fuse_attention: failed to match v path")
174 return
175
176 add_v, matmul_v = v_nodes[-2], v_nodes[-1]
177
178 causal_mask_input_index = None
179 add_mask = None
180 add_mask_indices = []
181 qk_nodes = self.model.match_parent_path(
182 matmul_qkv,
183 ["Softmax", "Reshape", "Add", "Reshape", "MatMul"],
184 [0, 0, 0, None, 0],
185 return_indice=add_mask_indices,
186 )
187 if qk_nodes is None:
188 qk_nodes = self.model.match_parent_path(
189 matmul_qkv,
190 ["Softmax", "MatMul"],
191 [0, 0],
192 )
193 if qk_nodes is None:
194 qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "Add", "Mul", "MatMul"], [0, 0, 0, 0])
195 if qk_nodes is not None:
196 add_mask = qk_nodes[1]
197 else:
198 # If attention mask is not used, we can still match the qk path.
199 qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "Mul", "MatMul"], [0, 0, 0])
200 if qk_nodes is None:
201 # Cast nodes are added in the model for fp16.
202 qk_nodes = self.model.match_parent_path(
203 matmul_qkv,
204 ["Cast", "Cast", "Softmax", "Add", "Mul", "MatMul"],
205 [0, 0, 0, 0, 0, 0],
206 )
207 if qk_nodes is not None:
208 add_mask = qk_nodes[3]
209 else:
210 # If attention mask is not used, we can still match the qk path.
211 qk_nodes = self.model.match_parent_path(
212 matmul_qkv,
213 ["Cast", "Cast", "Softmax", "Mul", "MatMul"],
214 [0, 0, 0, 0, 0],
215 )
216 if qk_nodes is None:
217 logger.debug("fuse_attention: failed to match qk path")
218 return
219 else:
220 assert len(add_mask_indices) == 1
221 causal_mask_input_index = 1 - add_mask_indices[0]
222 add_mask = qk_nodes[2]
223
224 matmul_qk = qk_nodes[-1]
225
226 q_nodes = self.model.match_parent_path(
227 matmul_qk,
228 ["Reshape", "Transpose", "Reshape", "Mul", "Add", "MatMul"],
229 [0, 0, 0, 0, None, None],
230 )
231 if q_nodes is None:
232 q_nodes = self.model.match_parent_path(
233 matmul_qk, ["Transpose", "Reshape", "Add", "MatMul"], [0, 0, 0, None]
234 )
235 if q_nodes is None:
236 logger.debug("fuse_attention: failed to match q path")
237 return
238
239 reshape_q = q_nodes[1]
240 else:
241 reshape_q = q_nodes[2]
242
243 add_q, matmul_q = q_nodes[-2], q_nodes[-1]
244
245 k_nodes = self.model.match_parent_path(
246 matmul_qk,
247 ["Transpose", "Reshape", "Transpose", "Reshape", "Add", "MatMul"],
248 [1, 0, 0, 0, 0, None],
249 )
250 if k_nodes is None:
251 k_nodes = self.model.match_parent_path(
252 matmul_qk, ["Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, None]
253 )
254 if k_nodes is None:
255 logger.debug("fuse_attention: failed to match k path")
256 return
257
258 add_k, matmul_k = k_nodes[-2], k_nodes[-1]
259
260 if matmul_q.input[0] != root_input or matmul_k.input[0] != root_input or matmul_v.input[0] != root_input:
261 logger.debug("fuse_attention: expect to have same input to q, k and v matmul")
262 return
263
264 num_heads, hidden_size = self.get_num_heads_and_hidden_size(reshape_q)
265 if num_heads <= 0 or hidden_size <= 0:
266 logger.debug("fuse_attention: failed to detect num_heads or hidden_size")
267 return
268
269 attention_last_node = reshape_qkv
270
271 add_qk = ""
272 causal_mask_nodes_1 = None
273 causal_mask_nodes_2 = None
274 if add_mask is not None:
275 if add_mask.input[1] == "attention_mask":
276 add_qk = add_mask.input[1]
277 else:
278 # 4D Add after Q x K'
279 add_qk_nodes = self.model.match_parent_path(
280 add_mask,
281 [
282 "Where",
283 "Sub",
284 "Cast",
285 "Expand",
286 "Unsqueeze",
287 "Unsqueeze",
288 "Reshape",
289 "Reshape",
290 "Cast",
291 ],
292 [1, 2, 1, 0, 0, 0, 0, 0, 0],
293 )
294 if add_qk_nodes is not None:
295 add_qk = add_mask.input[1]
296 else:
297 # Here we do not match the whole subgraph since it is very complex. Instead, we just check whether a key path
298 # of computing causal mask.
299 causal_mask_nodes_1 = self.model.match_parent_path(
300 add_mask,
301 ["Concat", "Expand", "Unsqueeze", "Unsqueeze", "Where", "Less"],
302 [causal_mask_input_index, 0, 0, 0, 0, 0],
303 )
304 # If the model is exported with batch_size == 1, there is no Concat node
305 causal_mask_nodes_2 = self.model.match_parent_path(
306 add_mask,
307 ["Expand", "Unsqueeze", "Unsqueeze", "Where", "Less"],
308 [causal_mask_input_index, 0, 0, 0, 0],
309 )
310
311 if causal_mask_nodes_1 is None and causal_mask_nodes_2 is None:
312 logger.debug("fuse_attention: failed to match causal mask subgraph")
313 return
314
315 new_node = self.create_attention_node(
316 mask_index=None,
317 q_matmul=matmul_q,
318 k_matmul=matmul_k,
319 v_matmul=matmul_v,
320 q_add=add_q,
321 k_add=add_k,
322 v_add=add_v,
323 num_heads=num_heads,
324 hidden_size=hidden_size,
325 first_input=root_input,
326 output=attention_last_node.output[0],
327 add_qk_str=add_qk,
328 scale=None,
329 causal=(causal_mask_nodes_1 is not None) or (causal_mask_nodes_2 is not None),
330 )
331 if new_node is None:
332 logger.debug("fuse_attention: failed to create fused node")
333 return
334
335 self.nodes_to_add.append(new_node)
336 self.node_name_to_graph_name[new_node.name] = self.this_graph_name
337 self.nodes_to_remove.extend([attention_last_node, transpose_qkv])
338
339 # Use prune graph to remove nodes since they are shared by all attention nodes.
340 self.prune_graph = True
341 