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_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionConstantFold(Fusion):
16 def __init__(self, model: OnnxModel):
17 super().__init__(model, "", ["Transpose"])
18 self.count = 0
19
20 def apply(self):
21 super().apply()
22 if self.count > 0:
23 logger.info(f"Constant Folded: {self.count}")
24
25 def fuse(self, node, input_name_to_nodes, output_name_to_node):
26 """
27 Apply multiple fusions on Transpose nodes that can be constant folded.
28 """
29 self.fuse_1(node, input_name_to_nodes, output_name_to_node)
30 self.fuse_2(node, input_name_to_nodes, output_name_to_node)
31
32 def fuse_1(self, node, input_name_to_nodes, output_name_to_node):
33 """
34 Constant fold any initializer data representing a MatMul's
35 weights that are stored in a Transpose op
36
37 Ex: Transpose --> Gemm or Transpose --> MatMul
38 """
39 # Check if Transpose node only has one input and one output
40 if len(node.input) != 1 or len(node.output) != 1:
41 logger.debug("fuse_constant_fold: node has more than one input or output")
42 return
43
44 # Check if input is initializer data
45 proto = self.model.get_initializer(node.input[0])
46 if proto is None:
47 logger.debug("fuse_constant_fold: failed to identify initializer input")
48 return
49
50 # Check that all nodes using input are Transpose ops that also only use the initializer data as input
51 skip = False
52 for child_node in input_name_to_nodes[node.input[0]]:
53 if not (child_node.op_type == "Transpose" and len(node.input) == 1):
54 skip = True
55 break
56 if skip:
57 logger.debug("fuse_constant_fold: other non-Transpose nodes use the initializer")
58 return
59
60 # Check that all nodes using output are Gemm or MatMul ops
61 for child_node in input_name_to_nodes[node.output[0]]:
62 if not (child_node.op_type == "Gemm" or child_node.op_type == "MatMul"):
63 skip = True
64 break
65 if skip:
66 logger.debug("fuse_constant_fold: other non-Gemm and non-MatMul nodes use the transposed data")
67 return
68
69 # Check if initializer data is 2D
70 weight = NumpyHelper.to_array(proto)
71 if len(weight.shape) != 2:
72 logger.debug("fuse_constant_fold: shape of initializer data is not 2D")
73 return
74
75 # Remove old TensorProto and add new TensorProto while re-using same name
76 name = proto.name
77 dtype = proto.data_type
78 self.remove_initializer(proto)
79 self.add_initializer(
80 name=name,
81 data_type=dtype,
82 dims=[weight.shape[1], weight.shape[0]],
83 vals=weight.T,
84 )
85
86 # Update weights input to be the initializer name and not
87 # the output of the Transpose op
88 for child_node in input_name_to_nodes[node.output[0]]:
89 for i in range(len(child_node.input)):
90 if child_node.input[i] == node.output[0]:
91 child_node.input[i] = node.input[0]
92
93 if child_node.op_type == "Gemm" and (i == 0 or i == 1):
94 # Ensure that transA/transB is set to 0 in Gemm
95 key = "transA" if i == 0 else "transB"
96 for j, attr_key in enumerate(child_node.attribute):
97 if attr_key.name == key:
98 child_node.attribute[j].i = 0
99
100 # Add node to list of nodes to remove
101 self.nodes_to_remove.append(node)
102 self.count += 1
103
104 def fuse_2(self, node, input_name_to_nodes, output_name_to_node):
105 """
106 Constant fold any Transpose --> Transpose ops since the root input
107 is the final result
108
109 Ex: root_input --> Transpose --> Transpose --> next_node to root_input --> next_node
110 """
111 # Check if Transpose node only has one input and one output
112 if len(node.input) != 1 or len(node.output) != 1:
113 logger.debug("fuse_constant_fold: node has more than one input or output")
114 return
115
116 # Check if parent node is Transpose node with only one input and one output
117 parent_node = self.model.match_parent(node, "Transpose", 0)
118 if parent_node is None:
119 logger.debug("fuse_constant_fold: failed to identify parent Transpose node")
120 return
121 if len(parent_node.input) != 1 or len(parent_node.output) != 1:
122 logger.debug("fuse_constant_fold: parent node has more than one input or output")
123 return
124
125 node_perm = node.attribute[0].ints
126 parent_node_perm = parent_node.attribute[0].ints
127
128 if node_perm != parent_node_perm:
129 logger.debug("fuse_constant_fold: Transpose node permutations aren't identical")
130 return
131
132 # For nodes that use output of child Transpose node as an input,
133 # replace that input with root_input
134 root_input = parent_node.input[0]
135 output_nodes = input_name_to_nodes[node.output[0]]
136 for output_node in output_nodes:
137 for i, input_ in enumerate(output_node.input):
138 if input_ == node.output[0]:
139 output_node.input[i] = root_input
140
141 # Add node to list of nodes to remove
142 self.nodes_to_remove.append(node)
143 self.nodes_to_remove.append(parent_node)
144 self.count += 1
145 