codekingpro/portable-devtools
114k
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6import json
7from argparse import ArgumentParser
8
9import onnx
10from onnx import TensorProto, helper
11
12
13def graph_topological_sort(graph):
14 deps_count = [0] * len(graph.node) # dependency count of each node
15 deps_to_nodes = {} # input to node indice
16 sorted_nodes = [] # initialize sorted_nodes
17 for node_idx, node in enumerate(graph.node):
18 # CANNOT use len(node.input) directly because input can be optional
19 deps_count[node_idx] = sum(1 for _ in node.input if _)
20 if deps_count[node_idx] == 0: # Constant doesn't depend on any inputs
21 sorted_nodes.append(graph.node[node_idx])
22 continue
23
24 for input_name in node.input:
25 if input_name not in deps_to_nodes:
26 deps_to_nodes[input_name] = [node_idx]
27 else:
28 deps_to_nodes[input_name].append(node_idx)
29
30 # Note: this logic only applies to top level graph since a sub graph could use intializer from parent graph
31 initializer_names = [init.name for init in graph.initializer]
32 graph_input_names = [input.name for input in graph.input]
33 input_names = initializer_names + graph_input_names
34 input_names.sort()
35 prev_input_name = None
36 for input_name in input_names:
37 if prev_input_name == input_name:
38 continue
39
40 prev_input_name = input_name
41 if input_name in deps_to_nodes:
42 for node_idx in deps_to_nodes[input_name]:
43 deps_count[node_idx] = deps_count[node_idx] - 1
44 if deps_count[node_idx] == 0:
45 sorted_nodes.append(graph.node[node_idx])
46
47 start = 0
48 end = len(sorted_nodes)
49
50 while start < end:
51 for output in sorted_nodes[start].output:
52 if output in deps_to_nodes:
53 for node_idx in deps_to_nodes[output]:
54 deps_count[node_idx] = deps_count[node_idx] - 1
55 if deps_count[node_idx] == 0:
56 sorted_nodes.append(graph.node[node_idx])
57 end = end + 1
58 start = start + 1
59
60 assert end == len(graph.node), "Graph is not a DAG"
61 graph.ClearField("node")
62 graph.node.extend(sorted_nodes)
63
64
65class QnnTensorStruct:
66 def __init__(self):
67 self.name = ""
68 self.onnx_data_type = TensorProto.FLOAT
69 self.dim = []
70
71
72def qnn_data_type_to_onnx_data_type(qnn_data_type):
73 # QNN_DATATYPE_UFIXED_POINT_8 QNN_DATATYPE_UINT_8
74 if qnn_data_type == 0x0408 or qnn_data_type == 0x0108:
75 return TensorProto.UINT8
76 # QNN_DATATYPE_UFIXED_POINT_16 QNN_DATATYPE_UINT_16
77 elif qnn_data_type == 0x0416 or qnn_data_type == 0x0116:
78 return TensorProto.UINT16
79 # QNN_DATATYPE_UFIXED_POINT_32 QNN_DATATYPE_UINT_32
80 elif qnn_data_type == 0x0432 or qnn_data_type == 0x0132:
81 return TensorProto.UINT32
82 # QNN_DATATYPE_UINT_64
83 elif qnn_data_type == 0x0164:
84 return TensorProto.UINT64
85 # QNN_DATATYPE_FIXED_POINT_8 QNN_DATATYPE_INT_8
86 elif qnn_data_type == 0x0308 or qnn_data_type == 0x0008:
87 return TensorProto.INT8
88 # QNN_DATATYPE_FIXED_POINT_16 QNN_DATATYPE_INT_16
89 elif qnn_data_type == 0x0316 or qnn_data_type == 0x0016:
90 return TensorProto.INT16
91 # QNN_DATATYPE_FIXED_POINT_32 QNN_DATATYPE_INT_32
92 elif qnn_data_type == 0x0332 or qnn_data_type == 0x0032:
93 return TensorProto.INT32
94 # QNN_DATATYPE_INT_64
95 elif qnn_data_type == 0x0064:
96 return TensorProto.INT64
97 # QNN_DATATYPE_FLOAT_16
98 elif qnn_data_type == 0x0216:
99 return TensorProto.FLOAT16
100 # QNN_DATATYPE_FLOAT_32
101 elif qnn_data_type == 0x0232:
102 return TensorProto.FLOAT
103 # QNN_DATATYPE_BOOL_8
104 elif qnn_data_type == 0x0508:
105 return TensorProto.BOOL
106 else:
107 return TensorProto.UNDEFINED
108
109
110def parse_qnn_json_file(qnn_json_file_path, qnn_input_output_tensor_dic):
111 with open(qnn_json_file_path) as qnn_json_file:
112 qnn_json = json.load(qnn_json_file)
113 assert "graph" in qnn_json, "QNN converted json file not valid. Can't find graph."
114 assert "tensors" in qnn_json["graph"], "QNN converted json file not valid. Can't find tensors."
115 for qnn_tensor_name, qnn_tensor_attribute in qnn_json["graph"]["tensors"].items():
116 # type:0 - QNN input tensor, type:1 - QNN output tensor
117 assert (
118 "type" in qnn_tensor_attribute
119 and "data_type" in qnn_tensor_attribute
120 and "dims" in qnn_tensor_attribute
121 ), "QNN converted json file not valid. Can't find some keys from tensors"
122 if qnn_tensor_attribute["type"] == 0 or qnn_tensor_attribute["type"] == 1:
123 qnn_tensor = QnnTensorStruct()
124 qnn_tensor.name = qnn_tensor_name
125 qnn_tensor.onnx_data_type = qnn_data_type_to_onnx_data_type(qnn_tensor_attribute["data_type"])
126 qnn_tensor.dim = qnn_tensor_attribute["dims"]
127 qnn_input_output_tensor_dic[qnn_tensor_name] = qnn_tensor
128
129 assert len(qnn_input_output_tensor_dic) > 1, (
130 "Converted QNN model not valid. It should have at least 1 input & 1 output."
131 )
132
133
134def compare_onnx_shape_with_qnn_shape(onnx_dims, qnn_dims):
135 assert len(onnx_dims) == len(qnn_dims), "Onnx shape and Qnn shape has different rank."
136 return all(onnx_dims[i].dim_value == qnn_dims[i] for i in range(len(onnx_dims)))
137
138
139def gen_to_channel_first_perm(rank):
140 assert rank > 2, "Shape rank should >2 for the Transpose node."
141 perm = []
142 perm.append(0)
143 perm.append(rank - 1)
144 for i in range(1, rank - 1):
145 perm.append(i) # noqa: PERF402
146
147 return perm
148
149
150def gen_to_channel_last_perm(rank):
151 assert rank > 2, "Shape rank should >2 for the Transpose node."
152 perm = []
153 perm.append(0)
154 for i in range(2, rank):
155 perm.append(i) # noqa: PERF402
156 perm.append(1)
157
158 return perm
159
160
161# Onnxruntime QNN EP can support context binary file generated by QNN tool chain. However QNN generated context binary file
162# uses channel last data layout and 8 bits or 16 bits for input and output.
163# This script gets the QNN model input & output information from QNN converted model_net.json file, compare them with Onnx model
164# and inserts Cast, Transpose nodes to Onnx model if required
165def main():
166 parser = ArgumentParser(
167 "Insert Cast, Transpose nodes into Onnx model to make it aligned with QNN generated context binary."
168 )
169 parser.add_argument("-m", "--onnx_model", help="Required. Path to Onnx model file.", required=True, type=str)
170 parser.add_argument(
171 "-q", "--qnn_json", help="Required. Path to Qnn converted model_net.json file.", required=True, type=str
172 )
173 args = parser.parse_args()
174
175 # Parse Qnn model_net.json file to get the graph input output information
176 qnn_input_output_tensor_dic = {}
177 parse_qnn_json_file(args.qnn_json, qnn_input_output_tensor_dic)
178
179 model = onnx.load(args.onnx_model)
180
181 nodes_to_add = []
182 # Tranch the tensor name change to update the consumer nodes
183 graph_input_output_name_dic = {}
184 for graph_input in model.graph.input:
185 if graph_input.name in qnn_input_output_tensor_dic:
186 input_name_fater_node_insert = graph_input.name
187 qnn_input_tensor = qnn_input_output_tensor_dic[graph_input.name]
188 # Insert Cast node if Onnx input and Qnn input has different data type
189 if graph_input.type.tensor_type.elem_type != qnn_input_tensor.onnx_data_type:
190 # Insert Cast node
191 cast_input_name = input_name_fater_node_insert
192 cast_output_name = cast_input_name + "_qnn_cast"
193 input_cast_node = helper.make_node(
194 "Cast",
195 name=cast_output_name,
196 inputs=[cast_input_name],
197 outputs=[cast_output_name],
198 to=graph_input.type.tensor_type.elem_type,
199 )
200 # Change input data type to Qnn input data type
201 graph_input.type.tensor_type.elem_type = qnn_input_tensor.onnx_data_type
202 nodes_to_add.extend([input_cast_node])
203 input_name_fater_node_insert = cast_output_name
204 graph_input_output_name_dic[graph_input.name] = cast_output_name
205
206 if not compare_onnx_shape_with_qnn_shape(graph_input.type.tensor_type.shape.dim, qnn_input_tensor.dim):
207 # Add Transpose node (channel last to channel first)
208 transpose_perm = gen_to_channel_first_perm(len(graph_input.type.tensor_type.shape.dim))
209 transpose_input_name = input_name_fater_node_insert
210 transpose_output_name = transpose_input_name + "_qnn_trans"
211 input_transpose_node = helper.make_node(
212 "Transpose",
213 name=transpose_output_name,
214 inputs=[transpose_input_name],
215 outputs=[transpose_output_name],
216 perm=transpose_perm,
217 )
218 nodes_to_add.extend([input_transpose_node])
219 graph_input_output_name_dic[graph_input.name] = transpose_output_name
220
221 # Change input shape to Qnn input shape
222 for i in range(len(graph_input.type.tensor_type.shape.dim)):
223 graph_input.type.tensor_type.shape.dim[i].dim_value = qnn_input_tensor.dim[i]
224 else:
225 raise AssertionError("Error: Onnx model input: " + graph_input.name + " not exist from QNN model input.")
226
227 for graph_output in model.graph.output:
228 if graph_output.name in qnn_input_output_tensor_dic:
229 output_name_after_node_insert = graph_output.name
230 # Insert Cast node if Onnx input and Qnn input has idfferent data type
231 qnn_output_tensor = qnn_input_output_tensor_dic[graph_output.name]
232 if graph_output.type.tensor_type.elem_type != qnn_output_tensor.onnx_data_type:
233 # Insert Cast node
234 cast_output_name = output_name_after_node_insert
235 cast_input_name = cast_output_name + "_qnn_cast"
236 output_cast_node = helper.make_node(
237 "Cast",
238 name=cast_input_name,
239 inputs=[cast_input_name],
240 outputs=[cast_output_name],
241 to=qnn_output_tensor.onnx_data_type,
242 )
243 # Change output data type to Onn output data type
244 graph_output.type.tensor_type.elem_type = qnn_output_tensor.onnx_data_type
245 nodes_to_add.extend([output_cast_node])
246 output_name_after_node_insert = cast_input_name
247 graph_input_output_name_dic[graph_output.name] = cast_input_name
248
249 if not compare_onnx_shape_with_qnn_shape(graph_output.type.tensor_type.shape.dim, qnn_output_tensor.dim):
250 # Add Transpose node (channel first to channel last)
251 transpose_perm = gen_to_channel_last_perm(len(graph_output.type.tensor_type.shape.dim))
252 transpose_output_name = output_name_after_node_insert
253 transpose_input_name = transpose_output_name + "_qnn_trans"
254 output_transpose_node = helper.make_node(
255 "Transpose",
256 name=transpose_input_name,
257 inputs=[transpose_input_name],
258 outputs=[transpose_output_name],
259 perm=transpose_perm,
260 )
261 nodes_to_add.extend([output_transpose_node])
262 graph_input_output_name_dic[graph_output.name] = transpose_input_name
263
264 # Change output shape to Qnn output shape
265 for i in range(len(graph_output.type.tensor_type.shape.dim)):
266 graph_output.type.tensor_type.shape.dim[i].dim_value = qnn_input_output_tensor_dic[
267 graph_output.name
268 ].dim[i]
269 else:
270 raise AssertionError("Error: Onnx model output: " + graph_output.name + " not exist from QNN model output.")
271
272 for node in model.graph.node:
273 for node_input_index, node_input in enumerate(node.input):
274 # update consumer node for graph inputs to connect to inserted node
275 if node_input in graph_input_output_name_dic:
276 node.input[node_input_index] = graph_input_output_name_dic[node_input]
277
278 for node_output_index, node_output in enumerate(node.output):
279 # update producer node for graph outputs to connect to inserted node
280 if node_output in graph_input_output_name_dic:
281 node.output[node_output_index] = graph_input_output_name_dic[node_output]
282
283 model.graph.node.extend(nodes_to_add)
284 graph_topological_sort(model.graph)
285
286 # Add extra parameter all_tensors_to_one_file=False, size_threshold=5000 if the model exceeds protobuf 2GB limit e.g below
287 # onnx.save(model, args.onnx_model.replace(".onnx", "_add_trans.onnx"), all_tensors_to_one_file=False, size_threshold=5000)
288 onnx.save(model, args.onnx_model.replace(".onnx", "_add_trans.onnx"))
289
290
291if __name__ == "__main__":
292 main()
293 