Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_gemmfastgelu.py122 linesDownload Raw Back to transformers
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 
codekingpro/portable-devtools · Team Ai