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 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 