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 onnx import NodeProto, TensorProto, helper, numpy_helper
10from onnx_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionAttentionVae(Fusion):
16 """
17 Fuse Attention subgraph of Vae Decoder into one Attention node.
18 """
19
20 def __init__(self, model: OnnxModel, hidden_size: int, num_heads: int):
21 super().__init__(model, "Attention", ["Softmax"])
22 self.hidden_size = hidden_size
23 self.num_heads = num_heads
24
25 # Flags to show warning only once
26 self.num_heads_warning = True
27 self.hidden_size_warning = True
28
29 def get_num_heads_and_hidden_size(self, reshape_q: NodeProto, add_q: NodeProto) -> tuple[int, int]:
30 """Detect num_heads and hidden_size from a reshape node.
31
32 Args:
33 reshape_q (NodeProto): reshape node for Q
34 add_q (NodeProto): add node for Q
35
36 Returns:
37 Tuple[int, int]: num_heads and hidden_size
38 """
39 concat = self.model.get_parent(reshape_q, 1)
40 if concat is None or len(concat.input) != 4:
41 return self.num_heads, self.hidden_size # Fall back to user specified value
42
43 value = self.model.get_constant_value(concat.input[2])
44 if not (value is not None and isinstance(value, np.ndarray) and value.size == 1):
45 return self.num_heads, self.hidden_size # Fall back to user specified value
46 num_heads = int(value)
47 if num_heads <= 0:
48 return self.num_heads, self.hidden_size # Fall back to user specified value
49
50 _, bias = self.model.get_constant_input(add_q)
51 if (bias is None) or (not isinstance(bias, np.ndarray)) or bias.ndim != 1:
52 return self.num_heads, self.hidden_size # Fall back to user specified value
53
54 hidden_size = bias.shape[0]
55
56 if self.num_heads > 0 and num_heads != self.num_heads:
57 if self.num_heads_warning:
58 logger.warning(
59 "Detected number of attention heads is %d. Ignore --num_heads %d", num_heads, self.num_heads
60 )
61 self.num_heads_warning = False # Do not show the warning more than once
62
63 if self.hidden_size > 0 and hidden_size != self.hidden_size:
64 if self.hidden_size_warning:
65 logger.warning("Detected hidden size is %d. Ignore --hidden_size %d", hidden_size, self.hidden_size)
66 self.hidden_size_warning = False # Do not show the warning more than once
67
68 return num_heads, hidden_size
69
70 def create_attention_node(
71 self,
72 q_matmul: NodeProto,
73 q_add: NodeProto,
74 k_matmul: NodeProto,
75 k_add: NodeProto,
76 v_matmul: NodeProto,
77 v_add: NodeProto,
78 num_heads: int,
79 hidden_size: int,
80 input_name: str,
81 output_name: str,
82 ) -> NodeProto | None:
83 """Create an Attention node.
84
85 Args:
86 q_matmul (NodeProto): MatMul node in fully connection for Q
87 q_add (NodeProto): Add bias node in fully connection for Q
88 k_matmul (NodeProto): MatMul node in fully connection for K
89 k_add (NodeProto): Add bias node in fully connection for K
90 v_matmul (NodeProto): MatMul node in fully connection for V
91 v_add (NodeProto): Add bias node in fully connection for V
92 num_heads (int): number of attention heads. If a model is pruned, it is the number of heads after pruning.
93 hidden_size (int): hidden dimension. If a model is pruned, it is the hidden dimension after pruning.
94 input_name (str): input name
95 output_name (str): output name
96
97 Returns:
98 Union[NodeProto, None]: the node created or None if failed.
99 """
100 if q_matmul.input[0] != input_name or k_matmul.input[0] != input_name or v_matmul.input[0] != input_name:
101 logger.debug(
102 "For self attention, input hidden state for q and k/v shall be same. Got %s, %s, %s",
103 q_matmul.input[0],
104 k_matmul.input[0],
105 v_matmul.input[0],
106 )
107 return None
108
109 if hidden_size > 0 and (hidden_size % num_heads) != 0:
110 logger.debug("input hidden size %d is not a multiple of num of heads %d", hidden_size, num_heads)
111 return None
112
113 q_weight_tensor = self.model.get_initializer(q_matmul.input[1])
114 k_weight_tensor = self.model.get_initializer(k_matmul.input[1])
115 v_weight_tensor = self.model.get_initializer(v_matmul.input[1])
116 if not (q_weight_tensor and k_weight_tensor and v_weight_tensor):
117 return None
118
119 q_bias_tensor = self.model.get_initializer(q_add.input[1]) or self.model.get_initializer(q_add.input[0])
120 k_bias_tensor = self.model.get_initializer(k_add.input[1]) or self.model.get_initializer(k_add.input[0])
121 v_bias_tensor = self.model.get_initializer(v_add.input[1]) or self.model.get_initializer(v_add.input[0])
122
123 q_bias = numpy_helper.to_array(q_bias_tensor)
124 k_bias = numpy_helper.to_array(k_bias_tensor)
125 v_bias = numpy_helper.to_array(v_bias_tensor)
126
127 q_bias_shape = np.prod(q_bias.shape)
128 k_bias_shape = np.prod(k_bias.shape)
129 v_bias_shape = np.prod(v_bias.shape)
130
131 # Sometimes weights are stored in fp16
132 if q_weight_tensor.data_type == 10:
133 logger.debug("weights are in fp16. Please run fp16 conversion after optimization")
134 return None
135
136 q_weight = numpy_helper.to_array(q_weight_tensor)
137 k_weight = numpy_helper.to_array(k_weight_tensor)
138 v_weight = numpy_helper.to_array(v_weight_tensor)
139
140 # assert q and k have same shape as expected
141 if q_weight.shape != k_weight.shape or q_weight.shape != v_weight.shape:
142 return None
143
144 qw_in_size = q_weight.shape[0]
145 kw_in_size = k_weight.shape[0]
146 vw_in_size = v_weight.shape[0]
147
148 assert qw_in_size == kw_in_size and kw_in_size == vw_in_size
149
150 if hidden_size > 0 and hidden_size != qw_in_size:
151 raise ValueError(
152 f"Input hidden size ({hidden_size}) is not same as weight dimension of q,k,v ({qw_in_size}). "
153 "Please provide a correct input hidden size or pass in 0"
154 )
155
156 # All the matrices can have the same shape or q, k matrics can have the same shape with v being different
157 # For 2d weights, the shapes would be [in_size, out_size].
158 # For 3d weights, shape would be [in_size, a, b] where a*b = out_size
159 qw_out_size = np.prod(q_weight.shape[1:])
160
161 qkv_weight = np.stack((q_weight, k_weight, v_weight), axis=1)
162 qkv_weight_dim = 3 * int(qw_out_size)
163
164 attention_node_name = self.model.create_node_name("Attention")
165
166 assert q_bias_shape == k_bias_shape == v_bias_shape
167
168 qkv_bias_dim = 0
169 qkv_bias = np.stack((q_bias, k_bias, v_bias), axis=0)
170 qkv_bias_dim = 3 * q_bias_shape
171
172 self.add_initializer(
173 name=attention_node_name + "_qkv_weight",
174 data_type=TensorProto.FLOAT,
175 dims=[qw_in_size, qkv_weight_dim],
176 vals=qkv_weight,
177 )
178
179 # No bias, use zeros
180 qkv_bias = np.zeros([3, hidden_size], dtype=np.float32)
181 qkv_bias_dim = 3 * hidden_size
182
183 self.add_initializer(
184 name=attention_node_name + "_qkv_bias",
185 data_type=TensorProto.FLOAT,
186 dims=[qkv_bias_dim],
187 vals=qkv_bias,
188 )
189
190 attention_inputs = [
191 input_name,
192 attention_node_name + "_qkv_weight",
193 attention_node_name + "_qkv_bias",
194 ]
195
196 attention_node = helper.make_node(
197 "Attention",
198 inputs=attention_inputs,
199 outputs=[output_name],
200 name=attention_node_name,
201 )
202 attention_node.domain = "com.microsoft"
203 attention_node.attribute.extend([helper.make_attribute("num_heads", num_heads)])
204
205 self.increase_counter("Attention (self attention)")
206 return attention_node
207
208 def fuse(self, softmax_node, input_name_to_nodes, output_name_to_node):
209 matmul_qkv = self.model.find_first_child_by_type(softmax_node, "MatMul", input_name_to_nodes, recursive=False)
210 if matmul_qkv is None:
211 return
212
213 reshape_qkv = self.model.find_first_child_by_type(matmul_qkv, "Reshape", input_name_to_nodes, recursive=False)
214 if reshape_qkv is None:
215 return
216
217 transpose_qkv = self.model.find_first_child_by_type(
218 reshape_qkv, "Transpose", input_name_to_nodes, recursive=False
219 )
220 if transpose_qkv is None:
221 return
222
223 reshape_out = self.model.find_first_child_by_type(
224 transpose_qkv, "Reshape", input_name_to_nodes, recursive=False
225 )
226 if reshape_out is None:
227 return
228
229 matmul_out = self.model.find_first_child_by_type(reshape_out, "MatMul", input_name_to_nodes, recursive=False)
230 if matmul_out is None:
231 return
232
233 add_out = self.model.find_first_child_by_type(matmul_out, "Add", input_name_to_nodes, recursive=False)
234 if add_out is None:
235 return
236
237 transpose_out = self.model.find_first_child_by_type(add_out, "Transpose", input_name_to_nodes, recursive=False)
238 if transpose_out is None:
239 return
240
241 v_nodes = self.model.match_parent_path(
242 matmul_qkv, ["Reshape", "Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, 0, None]
243 )
244 if v_nodes is None:
245 logger.debug("fuse_attention: failed to match v path")
246 return
247 (_, _, _, add_v, matmul_v) = v_nodes
248
249 qk_nodes = self.model.match_parent_path(matmul_qkv, ["Softmax", "Add", "Mul", "MatMul"], [0, 0, 0, 0])
250 if qk_nodes is not None:
251 (_softmax_qk, _add_zero, _mul_qk, matmul_qk) = qk_nodes
252 else:
253 logger.debug("fuse_attention: failed to match qk path")
254 return
255
256 q_nodes = self.model.match_parent_path(
257 matmul_qk, ["Reshape", "Transpose", "Reshape", "Add", "MatMul"], [0, 0, 0, 0, None]
258 )
259 if q_nodes is None:
260 logger.debug("fuse_attention: failed to match q path")
261 return
262 (_, _transpose_q, reshape_q, add_q, matmul_q) = q_nodes
263 k_nodes = self.model.match_parent_path(
264 matmul_qk, ["Transpose", "Reshape", "Transpose", "Reshape", "Add", "MatMul"], [1, 0, 0, 0, 0, None]
265 )
266 if k_nodes is None:
267 logger.debug("fuse_attention: failed to match k path")
268 return
269 (_, _, _, _, add_k, matmul_k) = k_nodes
270
271 attention_last_node = reshape_out
272
273 q_num_heads, q_hidden_size = self.get_num_heads_and_hidden_size(reshape_q, add_q)
274 if q_num_heads <= 0:
275 logger.debug("fuse_attention: failed to detect num_heads")
276 return
277
278 # number of heads are same for all the paths, hence to create attention node, we pass the q_num_heads
279 new_node = self.create_attention_node(
280 matmul_q,
281 add_q,
282 matmul_k,
283 add_k,
284 matmul_v,
285 add_v,
286 q_num_heads,
287 q_hidden_size,
288 matmul_q.input[0],
289 attention_last_node.output[0],
290 )
291 if new_node is None:
292 return
293
294 self.nodes_to_add.append(new_node)
295 self.node_name_to_graph_name[new_node.name] = self.this_graph_name
296
297 self.nodes_to_remove.extend([attention_last_node, transpose_qkv])
298
299 # Use prune graph to remove nodes since they are shared by all attention nodes.
300 self.prune_graph = True
301 