Team Ai
Datasetpublic

codekingpro/portable-devtools

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