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
8from fusion_base import Fusion
9from fusion_utils import NumpyHelper
10from onnx import NodeProto, TensorProto, helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionGemmFastGelu(Fusion):
17 def __init__(self, model: OnnxModel):
18 super().__init__(model, "GemmFastGelu", "FastGelu", "GemmFastGelu")
19 self.shape_infer = None
20 self.shape_infer_done = False
21
22 def get_dimensions_from_tensor_proto(self, tensor_proto: TensorProto) -> int | None:
23 if tensor_proto.type.tensor_type.HasField("shape"):
24 return len(tensor_proto.type.tensor_type.shape.dim)
25 else:
26 return None
27
28 def get_dimensions(self, input_name: str) -> int | None:
29 graph_input = self.model.find_graph_input(input_name)
30 if graph_input:
31 return self.get_dimensions_from_tensor_proto(graph_input)
32
33 if not self.shape_infer_done:
34 self.shape_infer = self.model.infer_runtime_shape(update=True)
35 self.shape_infer_done = True
36
37 if self.shape_infer is not None:
38 return self.get_dimensions_from_tensor_proto(self.shape_infer.known_vi_[input_name])
39
40 return None
41
42 def fuse(
43 self,
44 node: NodeProto,
45 input_name_to_nodes: dict[str, list[NodeProto]],
46 output_name_to_node: dict[str, NodeProto],
47 ):
48 """
49 This pattern is from PyTorch bert model
50 Fuse MatMul with FastGelu into one node:
51
52 [root] --> MatMul --> FastGelu -->
53
54 """
55 has_bias = False
56 if len(node.input) == 2:
57 has_bias = True
58
59 match_nodes = self.model.match_parent_path(node, ["MatMul"], [0])
60 if match_nodes is None:
61 return
62 matmul = match_nodes[0]
63
64 # matmul input X should >= two dimension, input weight should be two dimension
65 weight_index = -1
66 x_dims = 0
67 weight = None
68
69 for i, input in enumerate(matmul.input):
70 initializer = self.model.get_initializer(input)
71 if initializer is None:
72 x_dims = self.get_dimensions(matmul.input[i])
73 else:
74 weight_index = i
75 weight = NumpyHelper.to_array(initializer)
76 if weight is None:
77 return
78 if len(weight.shape) != 2:
79 return
80 if x_dims < len(weight.shape):
81 return
82
83 # bias weight should be one dimension
84 bias_index = -1
85 if has_bias:
86 bias_weight = None
87 for i, input in enumerate(node.input):
88 initializer = self.model.get_initializer(input)
89 if initializer is None:
90 continue
91 bias_index = i
92 bias_weight = NumpyHelper.to_array(initializer)
93 break
94 if bias_weight is None:
95 return
96 if len(bias_weight.shape) != 1:
97 return
98
99 subgraph_nodes = [node, matmul]
100 if not self.model.is_safe_to_fuse_nodes(
101 subgraph_nodes, [node.output[0]], input_name_to_nodes, output_name_to_node
102 ):
103 return
104
105 self.nodes_to_remove.extend(subgraph_nodes)
106
107 inputs = (
108 [matmul.input[1 - weight_index], matmul.input[weight_index], node.input[bias_index]]
109 if has_bias
110 else [matmul.input[1 - weight_index], matmul.input[weight_index]]
111 )
112
113 fused_node = helper.make_node(
114 "GemmFastGelu",
115 inputs=inputs,
116 outputs=node.output,
117 name=self.model.create_node_name("GemmFastGelu"),
118 )
119 fused_node.domain = "com.microsoft"
120 self.nodes_to_add.append(fused_node)
121 self.node_name_to_graph_name[fused_node.name] = self.this_graph_name
122 