Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_transpose.py168 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 FusionUtils
10from onnx import NodeProto, TensorProto, helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionTranspose(Fusion):
17    def __init__(self, model: OnnxModel):
18        super().__init__(model, "Transpose", "Transpose")
19
20    def fuse(
21        self,
22        transpose_node: NodeProto,
23        input_name_to_nodes: dict[str, list[NodeProto]],
24        output_name_to_node: dict[str, NodeProto],
25    ):
26        """
27        Note that onnxruntime will do comprehensive transpose optimization after loading model.
28        The purpose of this fusion is to make graph clean before running onnxruntime.
29
30        Case 1:
31              (input)-->Transpose(perm=a)-->Transpose(perm=b)-->
32        After:
33              (input)-->Transpose(perm=a)-->  (this path can be removed if the output is not used anymore)
34                |
35                +----->Transpose(perm=a*b)-->
36
37        Case 2 (Cast has only one child):
38              (input)-->Transpose(perm=a)--> Cast -->Transpose(perm=b)-->
39        After:
40              (input)-->Transpose(perm=a)-->  (this path can be removed if the output is not used anymore)
41                |
42                +----->Cast --> Transpose(perm=a*b)-->
43        """
44        transpose_b = transpose_node
45        if transpose_b.input[0] not in output_name_to_node:
46            return
47
48        transpose_a = output_name_to_node[transpose_b.input[0]]
49        if transpose_a.op_type != "Cast":
50            cast_node = None
51        else:
52            cast_node = transpose_a
53
54            cast_children = self.model.get_children(cast_node, input_name_to_nodes)
55            if cast_children and len(cast_children) > 1:
56                return
57
58            if cast_node.input[0] not in output_name_to_node:
59                return
60
61            transpose_a = output_name_to_node[cast_node.input[0]]
62
63        if transpose_a.op_type != "Transpose":
64            return
65
66        permutation = OnnxModel.get_node_attribute(transpose_b, "perm")
67        assert isinstance(permutation, list)
68
69        parent_permutation = OnnxModel.get_node_attribute(transpose_a, "perm")
70        assert isinstance(parent_permutation, list)
71
72        assert len(parent_permutation) == len(permutation)
73
74        output_permutation = []
75        for _j, index in enumerate(permutation):
76            output_permutation.append(parent_permutation[index])
77
78        if cast_node is None:
79            if FusionUtils.skip_parent(self.model, transpose_b, transpose_a, input_name_to_nodes):
80                self.nodes_to_remove.append(transpose_a)
81        else:
82            if FusionUtils.skip_parent(self.model, cast_node, transpose_a, input_name_to_nodes):
83                self.nodes_to_remove.append(transpose_a)
84        transpose_b.ClearField("attribute")
85        transpose_b.attribute.extend([helper.make_attribute("perm", output_permutation)])
86
87
88class FusionInsertTranspose(Fusion):
89    def __init__(self, model: OnnxModel):
90        super().__init__(model, "", "GroupNorm")
91
92    def create_transpose_node(self, input_name: str, perm: list[int], output_name=None):
93        """Append a Transpose node after an input"""
94        node_name = self.model.create_node_name("Transpose")
95        if output_name is None:
96            output_name = node_name + "_out" + "-" + input_name
97        transpose_node = helper.make_node("Transpose", inputs=[input_name], outputs=[output_name], name=node_name)
98        transpose_node.attribute.extend([helper.make_attribute("perm", perm)])
99        return transpose_node
100
101    def fuse(
102        self,
103        group_norm_node: NodeProto,
104        input_name_to_nodes: dict[str, list[NodeProto]],
105        output_name_to_node: dict[str, NodeProto],
106    ):
107        """
108        This optimization will insert an Transpose, and onnxruntime transpose optimizer will remove it together with
109        another Transpose so that we can get effect of reducing one Transpose after onnxruntime optimization.
110        Before:
111            --> Gemm --> Unsqueeze(axes=[2]) --> Unsqueeze(axes=[3]) --> Add --> Transpose([0,2,3,1]) --> GroupNorm
112        After:
113            --> Gemm --> Unsqueeze(axes=[1]) --> Unsqueeze(axes=[2]) -->Transpose([0,3,1,2]) --> Add --> Transpose([0,2,3,1]) --> GroupNorm
114        """
115        gemm_path = self.model.match_parent_path(
116            group_norm_node, ["Transpose", "Add", "Unsqueeze", "Unsqueeze", "Gemm"], [0, 0, None, 0, 0]
117        )
118        if gemm_path is None:
119            return
120        transpose, add, unsqueeze_3, unsqueeze_2, gemm = gemm_path
121        if self.model.find_graph_output(unsqueeze_3.output[0]):
122            return
123
124        permutation = OnnxModel.get_node_attribute(transpose, "perm")
125        assert isinstance(permutation, list)
126        if permutation != [0, 2, 3, 1]:
127            return
128
129        if not (
130            len(unsqueeze_3.input) == 2
131            and self.model.get_constant_value(unsqueeze_3.input[1]) == 3
132            and len(unsqueeze_2.input) == 2
133            and self.model.get_constant_value(unsqueeze_2.input[1]) == 2
134            and len(self.model.get_children(gemm, input_name_to_nodes)) == 1
135            and len(self.model.get_children(unsqueeze_3, input_name_to_nodes)) == 1
136            and len(self.model.get_children(unsqueeze_2, input_name_to_nodes)) == 1
137        ):
138            return
139
140        # Here we use hard-coded name so that it could be shared for the whole model.
141        axes_1 = "ort_const_unsqueeze_axes_1"
142        if self.model.get_initializer(axes_1) is None:
143            self.add_initializer(
144                name=axes_1,
145                data_type=TensorProto.INT64,
146                dims=[1],
147                vals=[1],
148                raw=False,
149            )
150
151        axes_2 = "ort_const_unsqueeze_axes_2"
152        if self.model.get_initializer(axes_2) is None:
153            self.add_initializer(
154                name=axes_2,
155                data_type=TensorProto.INT64,
156                dims=[1],
157                vals=[2],
158                raw=False,
159            )
160
161        unsqueeze_3.input[1] = "ort_const_unsqueeze_axes_2"
162        unsqueeze_2.input[1] = "ort_const_unsqueeze_axes_1"
163        transpose_output_name = self.model.create_node_name("Transpose") + "_NCHW"
164        self.model.replace_input_of_all_nodes(unsqueeze_3.output[0], transpose_output_name)
165        new_transpose = self.create_transpose_node(unsqueeze_3.output[0], [0, 3, 1, 2], transpose_output_name)
166        self.model.add_node(new_transpose, self.this_graph_name)
167        self.increase_counter("Insert Transpose")
168 
codekingpro/portable-devtools · Team Ai