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 helper, numpy_helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionNhwcConv(Fusion):
17 """Convert Conv to NhwcConv"""
18
19 def __init__(self, model: OnnxModel, update_weight=False):
20 super().__init__(model, "NhwcConv", ["Conv"], "NhwcConv")
21 self.update_weight = update_weight
22 self.fusion_utils = FusionUtils(model)
23
24 def create_transpose_node(self, input_name: str, perm: list[int], output_name=None):
25 """Append a Transpose node after an input"""
26 node_name = self.model.create_node_name("Transpose")
27
28 if output_name is None:
29 output_name = node_name + "_out" + "-" + input_name
30
31 transpose_node = helper.make_node("Transpose", inputs=[input_name], outputs=[output_name], name=node_name)
32 transpose_node.attribute.extend([helper.make_attribute("perm", perm)])
33
34 return transpose_node
35
36 def fuse(self, conv, input_name_to_nodes, output_name_to_node):
37 # Add Transpose node to convert input from NCHW to NHWC
38 input_transpose_node = self.create_transpose_node(conv.input[0], [0, 2, 3, 1])
39
40 nhwc_conv_input = input_transpose_node.output[0]
41
42 # Create a tensor for transposed weights (already in NHWC format).
43 node_name = self.model.create_node_name("NhwcConv")
44
45 # Make sure the weights is 4D
46 weight_tensor = self.model.get_initializer(conv.input[1])
47 if weight_tensor is None:
48 return
49 weight = numpy_helper.to_array(weight_tensor)
50 if len(weight.shape) != 4:
51 return
52
53 dtype = self.model.get_dtype(nhwc_conv_input)
54 if not (dtype is not None and weight_tensor.data_type == dtype):
55 cast_node = self.fusion_utils.add_cast_node(
56 input_name=nhwc_conv_input,
57 to_type=weight_tensor.data_type,
58 output_name_to_node=output_name_to_node,
59 )
60 nhwc_conv_input = cast_node.output[0]
61
62 if self.update_weight:
63 # Transpose weights from NCHW to NHWC
64 weight = weight.transpose(0, 2, 3, 1)
65
66 weight_name = node_name + "_weight_NHWC"
67 self.add_initializer(
68 name=weight_name,
69 data_type=weight_tensor.data_type,
70 dims=list(weight.shape),
71 vals=weight,
72 )
73 weight_transpose_node = None
74 else:
75 weight_transpose_node = self.create_transpose_node(conv.input[1], [0, 2, 3, 1])
76 weight_name = weight_transpose_node.output[0]
77
78 nhwc_output_name = node_name + "_out" + "-" + conv.output[0]
79 nhwc_conv = helper.make_node(
80 "NhwcConv",
81 inputs=[nhwc_conv_input, weight_name, *conv.input[2:]],
82 outputs=[nhwc_output_name],
83 name=node_name + "-" + conv.name,
84 )
85 nhwc_conv.attribute.extend(conv.attribute)
86 nhwc_conv.domain = "com.microsoft"
87
88 output_transpose_node = self.create_transpose_node(nhwc_conv.output[0], [0, 3, 1, 2], conv.output[0])
89
90 self.nodes_to_remove.append(conv)
91
92 nodes_to_add = [input_transpose_node, nhwc_conv, output_transpose_node]
93 if weight_transpose_node:
94 nodes_to_add.append(weight_transpose_node)
95 for node in nodes_to_add:
96 self.node_name_to_graph_name[node.name] = self.this_graph_name
97 self.nodes_to_add.extend(nodes_to_add)
98
99 self.increase_counter("NhwcConv")
100 