codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6from logging import getLogger
7
8import numpy as np
9from dynamo_onnx_helper import DynamoOnnxHelper
10from fusion_base import Fusion
11from fusion_options import AttentionOpType, FusionOptions
12from fusion_skiplayernorm import FusionBiasSkipLayerNormalization, FusionSkipLayerNormalization
13from fusion_utils import NumpyHelper
14from onnx import ModelProto, NodeProto, TensorProto, helper, numpy_helper
15from onnx_model import OnnxModel
16
17logger = getLogger(__name__)
18
19
20class ProcessGemmWFunc:
21 def __call__(self, x):
22 return np.transpose(x, (1, 0))
23
24
25class ProcessMatMulQFunc:
26 def __call__(self, x):
27 return np.transpose(np.split(x, 3, 0)[0], (1, 0))
28
29
30class ProcessMatMulKFunc:
31 def __call__(self, x):
32 return np.transpose(np.split(x, 3, 0)[1], (1, 0))
33
34
35class ProcessMatMulVFunc:
36 def __call__(self, x):
37 return np.transpose(np.split(x, 3, 0)[2], (1, 0))
38
39
40class ProcessBiasQFunc:
41 def __call__(self, x):
42 x = np.split(x, 3, -1)[0]
43 return x
44
45
46class ProcessBiasKFunc:
47 def __call__(self, x):
48 x = np.split(x, 3, -1)[1]
49 return x
50
51
52class ProcessBiasVFunc:
53 def __call__(self, x):
54 x = np.split(x, 3, -1)[2]
55 return x
56
57
58class ProcessRotCacheFunc:
59 def __call__(self, x):
60 # half rotary embedding
61 assert len(x.shape) == 2
62 if x.shape[1] == 32:
63 return x[:, 0:16]
64 return x
65
66
67# TODO: move to a separate file
68class Fission(Fusion):
69 def __init__(
70 self,
71 model: OnnxModel,
72 nodes_to_find: list[str],
73 ):
74 super().__init__(model, "DONOTUSE", nodes_to_find)
75
76 def set_attention_op_type(self, attn_op_type: AttentionOpType):
77 self.attn_op_type = attn_op_type
78
79 def get_uname(self, layer_id, name):
80 return name + "_" + str(layer_id)
81
82 def get_edge_by_name(self, edges, name):
83 for edge in edges:
84 if edge == name or edge.endswith(name) or edge.startswith(name):
85 return edge
86 raise ValueError(f"Edge {name} not found")
87
88 def get_input_by_name(self, node, name):
89 return self.get_edge_by_name(node.input, name)
90
91 def get_output_by_name(self, node, name):
92 return self.get_edge_by_name(node.output, name)
93
94 def process_initializer(self, initializer_name, functor, custom_name=None):
95 i = self.model.get_initializer(initializer_name)
96 i_np_array = NumpyHelper.to_array(i)
97 processed_i_np_array = functor(i_np_array)
98 new_tensor = helper.make_tensor(
99 initializer_name + "_processed" if custom_name is None else custom_name,
100 data_type=TensorProto.FLOAT,
101 dims=processed_i_np_array.shape,
102 vals=processed_i_np_array.flatten().tobytes(),
103 raw=True,
104 )
105 self.model.add_initializer(new_tensor, self.this_graph_name)
106 return new_tensor.name
107
108 def add_fp32_value_info(self, name):
109 new_value_info = self.model.graph().value_info.add()
110 new_value_info.name = name
111 new_value_info.type.tensor_type.elem_type = TensorProto.FLOAT
112
113 def add_int64_value_info(self, name):
114 new_value_info = self.model.graph().value_info.add()
115 new_value_info.name = name
116 new_value_info.type.tensor_type.elem_type = TensorProto.INT64
117
118 def replace_fp32_value_info(self, name, shape):
119 for value_info in self.model.graph().value_info:
120 if value_info.name == name:
121 self.model.graph().value_info.remove(value_info)
122 break
123 new_value_info = helper.make_tensor_value_info(
124 name,
125 elem_type=TensorProto.FLOAT,
126 shape=shape,
127 )
128 self.model.graph().value_info.extend([new_value_info])
129
130 def set_unique_name_and_add_nodes(
131 self, subgraph_nodes: list[NodeProto], layer_id: int, layer_known_edges_names: list[str]
132 ):
133 for new_node in subgraph_nodes:
134 for i, name in enumerate(new_node.input):
135 if name == "":
136 continue
137 elif name not in layer_known_edges_names:
138 new_node.input[i] = self.get_uname(layer_id, name)
139 self.add_fp32_value_info(new_node.input[i])
140 for i, name in enumerate(new_node.output):
141 if name == "":
142 continue
143 elif name not in layer_known_edges_names:
144 new_node.output[i] = self.get_uname(layer_id, name)
145 self.add_fp32_value_info(new_node.output[i])
146 new_node.name = self.get_uname(layer_id, new_node.name)
147 self.nodes_to_add.append(new_node)
148 self.node_name_to_graph_name[new_node.name] = self.this_graph_name
149
150 def layernorm(self, inputs: list[str], outputs: list[str], prefix: str = ""):
151 assert len(inputs) == 3
152 assert len(outputs) == 1
153 node = helper.make_node(
154 "LayerNormalization",
155 inputs=inputs,
156 outputs=outputs,
157 name=prefix + "_LayerNormalization",
158 epsilon=9.999999747378752e-06,
159 )
160 return [node]
161
162 def gemm(self, inputs: list[str], outputs: list[str], prefix: str = ""):
163 assert len(inputs) == 3
164 assert len(outputs) == 1
165 matmul = helper.make_node(
166 "MatMul",
167 inputs=[inputs[0], inputs[1]],
168 outputs=[prefix + "matmul_out"],
169 name=prefix + "MatMul",
170 )
171 add = helper.make_node(
172 "Add",
173 inputs=[prefix + "matmul_out", inputs[2]],
174 outputs=outputs,
175 name=prefix + "Bias",
176 )
177 return [matmul, add]
178
179 def rotary(self, inputs: list[str], outputs: list[str], prefix: str = "", rot_dim=32, num_heads=32):
180 assert len(inputs) == 4
181 assert len(outputs) == 1
182 node = helper.make_node(
183 "RotaryEmbedding",
184 inputs=inputs,
185 outputs=outputs,
186 name=prefix + "RotaryEmbedding",
187 domain="com.microsoft",
188 rotary_embedding_dim=rot_dim,
189 num_heads=num_heads,
190 )
191 return [node]
192
193 def fastgelu(self, inputs: list[str], outputs: list[str], prefix: str = ""):
194 assert len(inputs) == 1
195 assert len(outputs) == 1
196 node = helper.make_node(
197 "FastGelu",
198 inputs=inputs,
199 outputs=outputs,
200 name=prefix + "FastGelu",
201 domain="com.microsoft",
202 )
203 return [node]
204
205 def add(self, inputs: list[str], outputs: list[str], prefix: str = ""):
206 assert len(inputs) == 2
207 assert len(outputs) == 1
208 node = helper.make_node(
209 "Add",
210 inputs=inputs,
211 outputs=outputs,
212 name=prefix + "Add",
213 )
214 return [node]
215
216 def mha(self, inputs: list[str], outputs: list[str], prefix: str = "", num_heads=32):
217 assert len(inputs) == 8
218 assert len(outputs) == 3
219 node = helper.make_node(
220 "MultiHeadAttention",
221 inputs=inputs,
222 outputs=outputs,
223 name=prefix + "MultiHeadAttention",
224 domain="com.microsoft",
225 num_heads=num_heads,
226 unidirectional=1,
227 )
228 return [node]
229
230 def gqa(self, inputs: list[str], outputs: list[str], prefix: str = "", num_heads=32):
231 assert len(inputs) == 7
232 assert len(outputs) == 3
233 node = helper.make_node(
234 "GroupQueryAttention",
235 inputs=inputs,
236 outputs=outputs,
237 name=prefix + "GroupQueryAttention",
238 domain="com.microsoft",
239 num_heads=num_heads,
240 kv_num_heads=num_heads,
241 )
242 return [node]
243
244 def attention(self, inputs: list[str], outputs: list[str], prefix: str = "", num_heads=32):
245 assert len(inputs) == 5
246 assert len(outputs) == 2
247 node = helper.make_node(
248 "Attention",
249 inputs=inputs,
250 outputs=outputs,
251 name=prefix + "Attention",
252 domain="com.microsoft",
253 num_heads=num_heads,
254 unidirectional=1,
255 do_rotary=1,
256 rotary_embedding_dim=32,
257 )
258 return [node]
259
260 def paged_attn(
261 self,
262 inputs: list[str],
263 outputs: list[str],
264 prefix: str = "",
265 num_heads=32,
266 head_size=80,
267 scale=0.11180339753627777,
268 ):
269 assert len(inputs) == 6
270 assert len(outputs) == 1
271 node = helper.make_node(
272 "PagedAttention",
273 inputs=inputs,
274 outputs=outputs,
275 name=prefix + "PagedAttention",
276 domain="vllm.ort.ext",
277 num_heads=num_heads,
278 num_kv_heads=num_heads,
279 head_size=head_size,
280 scale=scale,
281 )
282 return [node]
283
284
285class Phi2PreProcessor(DynamoOnnxHelper):
286 def __init__(self, model: ModelProto, num_heads: int, hidden_size: int):
287 super().__init__(model)
288 self.num_hidden_layers = 32
289 self.num_attention_heads = num_heads
290 self.hidden_size = hidden_size
291
292 self.func_name = "modeling_phi_PhiModel_model_1"
293
294 def get_phi2_edge_dict(self) -> dict:
295 edge_dict = {}
296 edge_dict["lm_head_1"] = "logits"
297 edge_dict["l_input_ids_"] = "input_ids"
298 edge_dict["key_states"] = "past_key_0"
299 edge_dict["value_states"] = "past_value_0"
300 for i in range(1, self.num_hidden_layers, 1):
301 edge_dict[f"key_states_{i}"] = f"past_key_{i}"
302 edge_dict[f"value_states_{i}"] = f"past_value_{i}"
303 edge_dict[f"model_layers_{i}_1"] = f"present_key_{i}"
304 edge_dict[f"model_layers_{i}_1_1"] = f"present_value_{i}"
305
306 outputs = [o.name for o in self.model.graph.output]
307 if "model_layers_0_1_1" in outputs and "model_layers_0_1_2" in outputs:
308 edge_dict["model_layers_0_1_1"] = "present_key_0"
309 edge_dict["model_layers_0_1_2"] = "present_value_0"
310 else:
311 assert "model_layers_0_1" in outputs and "model_layers_0_1_1" in outputs
312 edge_dict["model_layers_0_1"] = "present_key_0"
313 edge_dict["model_layers_0_1_1"] = "present_value_0"
314 return edge_dict
315
316 def simplify_phi2_op_type(self):
317 phi2_transformer_layer_name = "modeling_phi_PhiDecoderLayer_model_layers"
318 for node in self.model.graph.node:
319 index = node.op_type.find(phi2_transformer_layer_name)
320 if index != -1:
321 node.op_type = node.op_type[index:]
322
323 def process_graph_io(self, attn_op_type: AttentionOpType):
324 self.use_attn = attn_op_type == AttentionOpType.Attention
325 self.use_vllm = attn_op_type == AttentionOpType.PagedAttention
326 graph = self.model.graph
327 new_inputs = []
328 for vi in graph.input:
329 if "input_ids" in vi.name:
330 vi_iid = helper.make_tensor_value_info(
331 vi.name,
332 elem_type=TensorProto.INT32 if not self.use_vllm else TensorProto.INT64,
333 shape=["batch_size", "seq_len"],
334 )
335 vi_step = helper.make_tensor_value_info(
336 "step",
337 elem_type=TensorProto.INT64,
338 shape=[1],
339 )
340 vi_pid = helper.make_tensor_value_info(
341 "position_ids",
342 elem_type=TensorProto.INT64,
343 shape=["batch_size", "seq_len"],
344 )
345 vi_mask = helper.make_tensor_value_info(
346 "attention_mask",
347 elem_type=TensorProto.INT32,
348 shape=["batch_size", "seq_len"],
349 )
350 vi_meta = helper.make_tensor_value_info(
351 "input_metadata",
352 elem_type=TensorProto.INT64,
353 shape=[1],
354 )
355 (
356 new_inputs.extend([vi_iid, vi_step, vi_mask])
357 if not self.use_vllm
358 else new_inputs.extend([vi_iid, vi_pid, vi_meta])
359 )
360 if self.use_attn:
361 if "past_key" in vi.name:
362 vi_cache = helper.make_tensor_value_info(
363 vi.name.replace("past_key", "past"),
364 elem_type=vi.type.tensor_type.elem_type,
365 shape=[
366 2,
367 "batch_size",
368 self.num_attention_heads,
369 "past_seq_len",
370 self.hidden_size // self.num_attention_heads,
371 ],
372 )
373 new_inputs.extend([vi_cache])
374 elif self.use_vllm:
375 if "past_key" in vi.name:
376 vi_cache = helper.make_tensor_value_info(
377 vi.name,
378 elem_type=vi.type.tensor_type.elem_type,
379 shape=["num_blocks", "num_heads", "head_size_x", "block_size", "block_x"],
380 )
381 new_inputs.extend([vi_cache])
382 if "past_value" in vi.name:
383 vi_cache = helper.make_tensor_value_info(
384 vi.name,
385 elem_type=vi.type.tensor_type.elem_type,
386 shape=[
387 "num_blocks",
388 "num_heads",
389 "head_size",
390 "block_size",
391 ],
392 )
393 new_inputs.extend([vi_cache])
394 else:
395 if "past_key" in vi.name or "past_value" in vi.name:
396 vi_cache = helper.make_tensor_value_info(
397 vi.name,
398 elem_type=vi.type.tensor_type.elem_type,
399 shape=[
400 "batch_size",
401 self.num_attention_heads,
402 "past_seq_len",
403 self.hidden_size // self.num_attention_heads,
404 ],
405 )
406 new_inputs.extend([vi_cache])
407
408 graph.ClearField("input")
409 graph.input.extend(new_inputs)
410
411 new_outputs = []
412 for i, vi in enumerate(graph.output):
413 if i == 0:
414 new_outputs.extend([vi])
415 else:
416 if self.use_attn:
417 if "present_key" in vi.name:
418 vi_cache = helper.make_tensor_value_info(
419 vi.name.replace("present_key", "present"),
420 elem_type=vi.type.tensor_type.elem_type,
421 shape=[
422 2,
423 "batch_size",
424 self.num_attention_heads,
425 "total_seq_len",
426 self.hidden_size // self.num_attention_heads,
427 ],
428 )
429 new_outputs.extend([vi_cache])
430 elif self.use_vllm:
431 pass
432 else:
433 vi_cache = helper.make_tensor_value_info(
434 vi.name,
435 elem_type=vi.type.tensor_type.elem_type,
436 shape=[
437 "batch_size",
438 self.num_attention_heads,
439 "total_seq_len",
440 self.hidden_size // self.num_attention_heads,
441 ],
442 )
443 new_outputs.extend([vi_cache])
444
445 graph.ClearField("output")
446 graph.output.extend(new_outputs)
447
448 def preprocess_onnx(self, attn_op_type: AttentionOpType):
449 function_name = None
450 for func in self.model.functions:
451 if func.name.endswith(self.func_name):
452 function_name = func.name
453 break
454 assert function_name is not None
455 self.unroll_function(function_name)
456 self.update_edges(self.get_phi2_edge_dict())
457 self.simplify_phi2_op_type()
458 self.remove_dropout_layer()
459 if attn_op_type == AttentionOpType.PagedAttention:
460 self.remove_lm_head_layer()
461 self.process_graph_io(attn_op_type)
462
463
464class FissionTransformerEmbeddingPhi(Fission):
465 def __init__(
466 self,
467 model: OnnxModel,
468 ):
469 super().__init__(model, ["torch_nn_modules_sparse_Embedding_model_embed_tokens_1"])
470
471 def fuse(self, node, input_name_to_nodes, output_name_to_node):
472 logger.info("Optimizing %s...", node.name)
473
474 assert len(node.input) == 2
475 assert len(node.output) == 1
476
477 input = node.input[0]
478 output = node.output[0]
479
480 embedding = self.get_input_by_name(node, "embed_tokens.weight")
481
482 layer_known_edges_names = [input, output, embedding]
483
484 subgraph_nodes = [
485 helper.make_node(
486 "Gather",
487 inputs=[embedding, input],
488 outputs=[output],
489 name="Embedding_Gather",
490 ),
491 ]
492
493 self.set_unique_name_and_add_nodes(subgraph_nodes, 0, layer_known_edges_names)
494 self.nodes_to_remove.append(node)
495 self.prune_graph = True
496
497
498class FissionTransformerLayerNormPhi(Fission):
499 def __init__(
500 self,
501 model: OnnxModel,
502 ):
503 super().__init__(model, ["torch_nn_modules_normalization_LayerNorm_model_final_layernorm_1"])
504
505 def fuse(self, node, input_name_to_nodes, output_name_to_node):
506 logger.info("Optimizing %s...", node.name)
507
508 assert len(node.input) == 3
509 assert len(node.output) == 1
510
511 input = node.input[0]
512 output = node.output[0]
513
514 ln_weight = self.get_input_by_name(node, "final_layernorm.weight")
515 ln_bias = self.get_input_by_name(node, "final_layernorm.bias")
516
517 layer_known_edges_names = [input, output, ln_weight, ln_bias]
518
519 subgraph_nodes = []
520 subgraph_nodes.extend(self.layernorm([input, ln_weight, ln_bias], [output], "Final"))
521
522 self.set_unique_name_and_add_nodes(subgraph_nodes, 99, layer_known_edges_names)
523
524 self.replace_fp32_value_info(input, ["batch_size", "seq_len", "hidden_size"])
525 self.replace_fp32_value_info(output, ["batch_size", "seq_len", "hidden_size"])
526
527 self.nodes_to_remove.append(node)
528 self.prune_graph = True
529
530
531class FissionTransformerCausalLMHeadPhi(Fission):
532 def __init__(
533 self,
534 model: OnnxModel,
535 ):
536 super().__init__(model, ["torch_nn_modules_linear_Linear_lm_head_1"])
537
538 def fuse(self, node, input_name_to_nodes, output_name_to_node):
539 logger.info("Optimizing %s...", node.name)
540
541 assert len(node.input) == 5
542 assert len(node.output) == 1
543
544 input = node.input[2]
545 output = node.output[0]
546
547 fc_weight = self.process_initializer(self.get_input_by_name(node, "lm_head.weight"), ProcessGemmWFunc())
548 fc_bias = self.get_input_by_name(node, "lm_head.bias")
549
550 layer_known_edges_names = [input, output, fc_weight, fc_bias]
551
552 subgraph_nodes = []
553 subgraph_nodes.extend(self.gemm([input, fc_weight, fc_bias], [output], "LMHead_"))
554
555 self.set_unique_name_and_add_nodes(subgraph_nodes, 99, layer_known_edges_names)
556
557 self.replace_fp32_value_info(input, ["batch_size", "seq_len", "hidden_size"])
558 self.replace_fp32_value_info(output, ["batch_size", "seq_len", 51200])
559
560 self.nodes_to_remove.append(node)
561 self.prune_graph = True
562
563
564class FissionTransformerBlockPhi(Fission):
565 def __init__(
566 self,
567 model: OnnxModel,
568 num_heads: int,
569 ):
570 self.num_heads = num_heads
571 max_num_layers = 32
572 self.func_to_layer_id = {}
573 nodes_to_find = []
574 for layer in range(max_num_layers):
575 func_name = f"modeling_phi_PhiDecoderLayer_model_layers_{layer}_1"
576 nodes_to_find.append(func_name)
577 self.func_to_layer_id[func_name] = layer
578
579 super().__init__(model, nodes_to_find)
580
581 def get_layer_id(self, node):
582 return self.func_to_layer_id[node.op_type]
583
584 def get_gqa_aux_nodes(self):
585 gqa_aux_nodes = [
586 helper.make_node(
587 "Cast",
588 inputs=["attention_mask"],
589 outputs=["mask_int64"],
590 name="Cast_gqa_aux_0",
591 to=TensorProto.INT64,
592 ),
593 helper.make_node(
594 "ReduceSum",
595 inputs=["mask_int64", "one"],
596 outputs=["mask_row_sums"],
597 name="ReduceSum_gqa_aux",
598 ),
599 helper.make_node(
600 "Sub",
601 inputs=["mask_row_sums", "one"],
602 outputs=["seqlens_k_int64"],
603 name="Sub_gqa_aux",
604 ),
605 helper.make_node(
606 "Cast",
607 inputs=["seqlens_k_int64"],
608 outputs=["seqlens_k"],
609 name="Cast_gqa_aux_1",
610 to=TensorProto.INT32,
611 ),
612 helper.make_node("Shape", inputs=["mask_int64"], outputs=["mask_shape"], name="Shape_gqa_aux_0"),
613 helper.make_node(
614 "Gather",
615 inputs=["mask_shape", "one"],
616 outputs=["total_seq_len_int64"],
617 name="Gather_gqa_aux_0",
618 axis=0,
619 ),
620 helper.make_node(
621 "Cast",
622 inputs=["total_seq_len_int64"],
623 outputs=["total_sequence_length"],
624 name="Cast_gqa_aux_2",
625 to=TensorProto.INT32,
626 ),
627 ]
628 return gqa_aux_nodes
629
630 def pack_qkv_gemm(self, q_w, k_w, v_w, q_b, k_b, v_b, weight_name, bias_name):
631 q_weight = self.model.get_initializer(q_w)
632 k_weight = self.model.get_initializer(k_w)
633 v_weight = self.model.get_initializer(v_w)
634 qw = np.transpose(NumpyHelper.to_array(q_weight), (1, 0))
635 kw = np.transpose(NumpyHelper.to_array(k_weight), (1, 0))
636 vw = np.transpose(NumpyHelper.to_array(v_weight), (1, 0))
637 qkv_weight = np.stack((qw, kw, vw), axis=1)
638
639 q_bias = self.model.get_initializer(q_b)
640 k_bias = self.model.get_initializer(k_b)
641 v_bias = self.model.get_initializer(v_b)
642 qb = NumpyHelper.to_array(q_bias)
643 kb = NumpyHelper.to_array(k_bias)
644 vb = NumpyHelper.to_array(v_bias)
645 qkv_bias = np.stack((qb, kb, vb), axis=0)
646
647 hidden_size = qkv_weight.shape[0]
648
649 weight = helper.make_tensor(
650 weight_name,
651 data_type=TensorProto.FLOAT,
652 dims=[hidden_size, hidden_size * 3],
653 vals=qkv_weight.flatten().tobytes(),
654 raw=True,
655 )
656 self.model.add_initializer(weight, self.this_graph_name)
657
658 bias = helper.make_tensor(
659 bias_name,
660 data_type=TensorProto.FLOAT,
661 dims=[hidden_size * 3],
662 vals=qkv_bias.flatten().tobytes(),
663 raw=True,
664 )
665 self.model.add_initializer(bias, self.this_graph_name)
666
667 self.add_fp32_value_info(weight.name)
668 self.add_fp32_value_info(bias.name)
669
670 return weight_name, bias_name
671
672 def fuse(
673 self,
674 node,
675 input_name_to_nodes,
676 output_name_to_node,
677 ):
678 logger.info("Optimizing %s...", node.name)
679
680 logger.info(f"AttentionOpType: {self.attn_op_type}")
681
682 layer_id = self.get_layer_id(node)
683
684 i_hidden_states = node.input[0]
685 i_key_cache = self.get_input_by_name(node, "past_key")
686 i_value_cache = self.get_input_by_name(node, "past_value")
687
688 o_hidden_states = node.output[-1]
689 o_key_cache = self.get_output_by_name(node, "present_key")
690 o_value_cache = self.get_output_by_name(node, "present_value")
691
692 ln_weight = self.get_input_by_name(node, "input_layernorm.weight")
693 ln_bias = self.get_input_by_name(node, "input_layernorm.bias")
694
695 attn_q_weight, attn_q_bias, attn_k_weight, attn_k_bias, attn_v_weight, attn_v_bias = (
696 None,
697 None,
698 None,
699 None,
700 None,
701 None,
702 )
703 attn_qkv_weight, attn_qkv_bias = None, None
704 cos_cache, sin_cache = None, None
705
706 if self.attn_op_type != AttentionOpType.Attention:
707 attn_q_weight = self.process_initializer(
708 self.get_input_by_name(node, "self_attn.q_proj.weight"), ProcessGemmWFunc()
709 )
710 attn_k_weight = self.process_initializer(
711 self.get_input_by_name(node, "self_attn.k_proj.weight"), ProcessGemmWFunc()
712 )
713 attn_v_weight = self.process_initializer(
714 self.get_input_by_name(node, "self_attn.v_proj.weight"), ProcessGemmWFunc()
715 )
716 attn_q_bias = self.get_input_by_name(node, "self_attn.q_proj.bias")
717 attn_k_bias = self.get_input_by_name(node, "self_attn.k_proj.bias")
718 attn_v_bias = self.get_input_by_name(node, "self_attn.v_proj.bias")
719
720 cos_cache = self.process_initializer(
721 self.get_input_by_name(node, "rotary_emb.cos_cached"), ProcessRotCacheFunc()
722 )
723 sin_cache = self.process_initializer(
724 self.get_input_by_name(node, "rotary_emb.sin_cached"), ProcessRotCacheFunc()
725 )
726 else:
727 attn_qkv_weight, attn_qkv_bias = self.pack_qkv_gemm(
728 self.get_input_by_name(node, "self_attn.q_proj.weight"),
729 self.get_input_by_name(node, "self_attn.k_proj.weight"),
730 self.get_input_by_name(node, "self_attn.v_proj.weight"),
731 self.get_input_by_name(node, "self_attn.q_proj.bias"),
732 self.get_input_by_name(node, "self_attn.k_proj.bias"),
733 self.get_input_by_name(node, "self_attn.v_proj.bias"),
734 self.get_uname(layer_id, "attn_qkv_weight"),
735 self.get_uname(layer_id, "attn_qkv_bias"),
736 )
737
738 attn_out_weight = self.process_initializer(
739 self.get_input_by_name(node, "self_attn.dense.weight"), ProcessGemmWFunc()
740 )
741 attn_out_bias = self.get_input_by_name(node, "self_attn.dense.bias")
742
743 mlp_fc1_weight = self.process_initializer(self.get_input_by_name(node, "mlp.fc1.weight"), ProcessGemmWFunc())
744 mlp_fc2_weight = self.process_initializer(self.get_input_by_name(node, "mlp.fc2.weight"), ProcessGemmWFunc())
745 mlp_fc1_bias = self.get_input_by_name(node, "mlp.fc1.bias")
746 mlp_fc2_bias = self.get_input_by_name(node, "mlp.fc2.bias")
747
748 layer_known_edges_names = []
749 layer_known_edges_names.extend([i_hidden_states, i_key_cache, i_value_cache])
750 layer_known_edges_names.extend([o_hidden_states, o_key_cache, o_value_cache])
751 layer_known_edges_names.extend([ln_weight, ln_bias])
752 if self.attn_op_type != AttentionOpType.Attention:
753 layer_known_edges_names.extend(
754 [
755 attn_q_weight,
756 attn_q_bias,
757 attn_k_weight,
758 attn_k_bias,
759 attn_v_weight,
760 attn_v_bias,
761 cos_cache,
762 sin_cache,
763 ]
764 )
765 else:
766 layer_known_edges_names.extend([attn_qkv_weight, attn_qkv_bias])
767 layer_known_edges_names.extend(
768 [attn_out_weight, attn_out_bias, mlp_fc1_weight, mlp_fc1_bias, mlp_fc2_weight, mlp_fc2_bias]
769 )
770 layer_known_edges_names.extend(
771 ["attention_mask", "step", "seqlens_k", "total_sequence_length", "input_metadata", "position_ids"]
772 )
773
774 subgraph_nodes = []
775 subgraph_nodes.extend(self.layernorm([i_hidden_states, ln_weight, ln_bias], ["ln_out"]))
776 subgraph_nodes.extend(self.gemm(["attn_out", attn_out_weight, attn_out_bias], ["attn_add_out"], "OutProj_"))
777 subgraph_nodes.extend(self.gemm(["ln_out", mlp_fc1_weight, mlp_fc1_bias], ["fc1_out"], "FC1_"))
778 subgraph_nodes.extend(self.fastgelu(["fc1_out"], ["gelu_out"]))
779 subgraph_nodes.extend(self.gemm(["gelu_out", mlp_fc2_weight, mlp_fc2_bias], ["fc2_out"], "FC2_"))
780 subgraph_nodes.extend(self.add(["attn_add_out", "fc2_out"], ["residual_1_out"], "Residual_1"))
781 subgraph_nodes.extend(self.add([i_hidden_states, "residual_1_out"], [o_hidden_states], "Residual_2"))
782 if self.attn_op_type != AttentionOpType.Attention:
783 subgraph_nodes.extend(self.gemm(["ln_out", attn_q_weight, attn_q_bias], ["query"], "Q_"))
784 subgraph_nodes.extend(self.gemm(["ln_out", attn_k_weight, attn_k_bias], ["key"], "K_"))
785 subgraph_nodes.extend(self.gemm(["ln_out", attn_v_weight, attn_v_bias], ["value"], "V_"))
786 # vllm engine requires full position ids as the input
787 pos_ids_name = "position_ids" if self.attn_op_type == AttentionOpType.PagedAttention else "step"
788 subgraph_nodes.extend(self.rotary(["query", pos_ids_name, cos_cache, sin_cache], ["query_rot"], "Q_"))
789 subgraph_nodes.extend(self.rotary(["key", pos_ids_name, cos_cache, sin_cache], ["key_rot"], "K_"))
790 if self.attn_op_type == AttentionOpType.MultiHeadAttention:
791 subgraph_nodes.extend(
792 self.mha(
793 ["query_rot", "key_rot", "value", "", "attention_mask", "", i_key_cache, i_value_cache],
794 ["attn_out", o_key_cache, o_value_cache],
795 )
796 )
797 elif self.attn_op_type == AttentionOpType.GroupQueryAttention:
798 subgraph_nodes.extend(
799 self.gqa(
800 [
801 "query_rot",
802 "key_rot",
803 "value",
804 i_key_cache,
805 i_value_cache,
806 "seqlens_k",
807 "total_sequence_length",
808 ],
809 ["attn_out", o_key_cache, o_value_cache],
810 )
811 )
812 if layer_id == 0:
813 gqa_aux_nodes = self.get_gqa_aux_nodes()
814 for new_node in gqa_aux_nodes:
815 self.nodes_to_add.append(new_node)
816 self.node_name_to_graph_name[new_node.name] = self.this_graph_name
817 self.model.add_initializer(
818 numpy_helper.from_array(np.array([1], dtype="int64"), name="one"), self.this_graph_name
819 )
820 elif self.attn_op_type == AttentionOpType.PagedAttention:
821 subgraph_nodes.extend(
822 self.paged_attn(
823 ["query_rot", "key_rot", "value", i_key_cache, i_value_cache, "input_metadata"],
824 ["attn_out"],
825 )
826 )
827 else:
828 past_name = f"past_{layer_id}"
829 present_name = f"present_{layer_id}"
830 layer_known_edges_names.extend([past_name, present_name])
831 subgraph_nodes.extend(
832 self.attention(
833 ["ln_out", attn_qkv_weight, attn_qkv_bias, "attention_mask", past_name], ["attn_out", present_name]
834 )
835 )
836
837 self.set_unique_name_and_add_nodes(subgraph_nodes, layer_id, layer_known_edges_names)
838
839 self.replace_fp32_value_info(i_hidden_states, ["batch_size", "seq_len", "hidden_size"])
840 self.replace_fp32_value_info(o_hidden_states, ["batch_size", "seq_len", "hidden_size"])
841
842 self.nodes_to_remove.append(node)
843 self.prune_graph = True
844
845
846class PhiOnnxModel(OnnxModel):
847 def __init__(self, model: ModelProto, num_heads: int, hidden_size: int):
848 super().__init__(model)
849 self.phi2_preprocessor = Phi2PreProcessor(self.model, num_heads, hidden_size)
850 self.fission_transformer_block = FissionTransformerBlockPhi(self, num_heads)
851 self.fission_causal_lm_head = FissionTransformerCausalLMHeadPhi(self)
852 self.fission_transformer_layernorm = FissionTransformerLayerNormPhi(self)
853 self.fission_transformer_embedding = FissionTransformerEmbeddingPhi(self)
854
855 def optimize(self, options: FusionOptions | None = None, add_dynamic_axes: bool = False):
856 assert options is not None
857 attn_op_type = options.attention_op_type
858
859 self.fission_transformer_block.set_attention_op_type(attn_op_type)
860
861 self.phi2_preprocessor.preprocess_onnx(attn_op_type)
862
863 self.fission_transformer_block.apply()
864 self.fission_transformer_layernorm.apply()
865 self.fission_causal_lm_head.apply()
866 self.fission_transformer_embedding.apply()
867
868 super().prune_graph()
869
870 # SLN ctor is placed here intentionally to delay the symbolic shape inference
871 self.fuse_sln = FusionSkipLayerNormalization(self)
872 self.fuse_bias_sln = FusionBiasSkipLayerNormalization(self)
873 self.fuse_sln.apply()
874 self.fuse_bias_sln.apply()
875
876 def get_fused_operator_statistics(self):
877 """
878 Returns node count of fused operators.
879 """
880 op_count = {}
881 ops = [
882 "Attention",
883 "MultiHeadAttention",
884 "GroupQueryAttention",
885 "PagedAttention",
886 "Gelu",
887 "BiasGelu",
888 "FastGelu",
889 "LayerNormalization",
890 "SkipLayerNormalization",
891 ]
892 for op in ops:
893 nodes = self.get_nodes_by_op_type(op)
894 op_count[op] = len(nodes)
895
896 logger.info(f"Optimized operators: {op_count}")
897 return op_count
898
899 def is_fully_optimized(self, fused_op_count=None):
900 """
901 Returns True when the model is fully optimized.
902 """
903 if fused_op_count is None:
904 fused_op_count = self.get_fused_operator_statistics()
905
906 def op_count(op_name: str):
907 return fused_op_count.get(op_name) or 0
908
909 attention = (
910 op_count("Attention")
911 + op_count("MultiHeadAttention")
912 + op_count("GroupQueryAttention")
913 + op_count("PagedAttention")
914 )
915 gelu = op_count("Gelu") + op_count("BiasGelu") + op_count("FastGelu")
916 layer_norm = op_count("LayerNormalization") + op_count("SkipLayerNormalization")
917
918 is_perfect = (attention > 0) and (attention == gelu) and (layer_norm >= attention)
919
920 if layer_norm == 0:
921 logger.debug("Layer Normalization not fused")
922
923 if gelu == 0:
924 logger.debug("Gelu (or FastGelu) not fused")
925
926 if attention == 0:
927 logger.warning("Attention (or MultiHeadAttention) not fused")
928
929 return is_perfect
930 