Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_skip_group_norm.py255 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5from logging import getLogger
6
7from fusion_base import Fusion
8from fusion_utils import NumpyHelper
9from onnx import helper
10from onnx_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionSkipGroupNorm(Fusion):
16    """
17    Fuse Add + GroupNorm into one node: SkipGroupNorm.
18    """
19
20    def __init__(self, model: OnnxModel):
21        super().__init__(model, "SkipGroupNorm", "GroupNorm")
22        # Update shape inference is needed since other fusions might add new edge which does not have shape info yet.
23        self.shape_infer_helper = self.model.infer_runtime_shape(update=True)
24
25        if self.shape_infer_helper is None:
26            logger.warning("SkipGroupNorm fusion will be skipped since symbolic shape inference disabled or failed.")
27
28    def create_transpose_node(self, input_name: str, perm: list[int], output_name=None):
29        """Append a Transpose node after an input"""
30        node_name = self.model.create_node_name("Transpose")
31        if output_name is None:
32            output_name = node_name + "_out" + "-" + input_name
33        transpose_node = helper.make_node("Transpose", inputs=[input_name], outputs=[output_name], name=node_name)
34        transpose_node.attribute.extend([helper.make_attribute("perm", perm)])
35        return transpose_node
36
37    def get_skip_index(self, add, is_channel_last: bool):
38        """Add has two inputs. This classifies which input is skip based on shape info (skip allows broadcast)."""
39        skip = -1
40        broadcast = False
41
42        assert self.shape_infer_helper is not None
43        shape_a = self.shape_infer_helper.get_edge_shape(add.input[0])
44        shape_b = self.shape_infer_helper.get_edge_shape(add.input[1])
45        assert shape_a is not None and shape_b is not None
46
47        if len(shape_a) == 4 and len(shape_b) == 4:
48            if shape_a == shape_b:
49                skip = 1
50            else:
51                c = 3 if is_channel_last else 1
52                h = 1 if is_channel_last else 2
53                w = 2 if is_channel_last else 3
54                if shape_a[0] == shape_b[0] and shape_a[c] == shape_b[c]:
55                    if shape_b[h] == 1 and shape_b[w] == 1:
56                        skip = 1
57                        broadcast = True
58                    elif shape_a[h] == 1 and shape_a[w] == 1:
59                        skip = 0
60                        broadcast = True
61
62        if skip < 0:
63            logger.debug(
64                "skip SkipGroupNorm fusion since shape of Add inputs (%s, %s) are not expected",
65                add.input[0],
66                add.input[1],
67            )
68        return skip, broadcast
69
70    def has_multiple_consumers(self, output_name, input_name_to_nodes):
71        """Whether an output has multiple consumers (like graph output or more than one children nodes)"""
72        return self.model.find_graph_output(output_name) is not None or (
73            output_name in input_name_to_nodes and len(input_name_to_nodes[output_name]) > 1
74        )
75
76    def remove_if_safe(self, node, input_name_to_nodes):
77        """Remove a node if it is safe (only one children, and not graph output)"""
78        if not self.has_multiple_consumers(node.output[0], input_name_to_nodes):
79            self.nodes_to_remove.extend([node])
80
81    def is_bias_1d(self, bias_name: str):
82        """Whether bias is an initializer of one dimension"""
83        initializer = self.model.get_initializer(bias_name)
84        if initializer is None:
85            return False
86
87        bias_weight = NumpyHelper.to_array(initializer)
88        if bias_weight is None:
89            logger.debug("Bias weight not found")
90            return False
91
92        if len(bias_weight.shape) != 1:
93            logger.debug("Bias weight is not 1D")
94            return False
95        return True
96
97    def match_bias_path(self, node, input_name_to_nodes, output_name_to_node):
98        """
99        Match the bias graph pattern from an Transpose node after Reshape node like in below example.
100        It checks whether the bias is 1D initializer. If so, remove Add and redirect MatMul output to Reshape.
101        """
102        # Before Fusion:
103        #                        MatMul  (bias)
104        #                            \  /     (shape)
105        #                             Add    /
106        #                               \   /
107        #       (a)                   Reshape
108        #        \                       |
109        # Transpose([0, 3, 1, 2])   Transpose([0, 3, 1, 2])  --- the start node, this func only handles the above nodes.
110        #                        \  /
111        #                         Add
112        #                         / \
113        #                      (c)  Transpose([0,2,3,1])
114        #                              |
115        #                           GroupNorm
116        #                              |
117        #                             (d)
118        #
119        # After Fusion (the nodes below Reshape is handled in the fuse function):
120        #                    MatMul (shape)
121        #                       \   /
122        #                (a)   Reshape
123        #                  \    /
124        #                SkipGroupNorm
125        #                  /    \
126        #                (d)   Transpose([0, 3, 1, 2])
127        #                         \
128        #                         (c)
129
130        add_input_index = []
131        bias_nodes = self.model.match_parent_path(
132            node, ["Reshape", "Add", "MatMul"], [0, 0, None], output_name_to_node, add_input_index
133        )
134        if bias_nodes is None:
135            return None
136
137        (reshape, add_bias, matmul) = bias_nodes
138        bias = bias_nodes[1].input[1 - add_input_index[0]]
139        if not self.is_bias_1d(bias):
140            return None
141
142        reshape.input[0] = matmul.output[0]
143        self.remove_if_safe(add_bias, input_name_to_nodes)
144
145        return bias
146
147    def match_transpose_from_nhwc(self, output_name, input_name_to_nodes, output_name_to_node):
148        """Match whether an output is from a Transpose(perm=[0,3,1,2]) node."""
149        parent = output_name_to_node.get(output_name, None)
150        if parent is not None and parent.op_type == "Transpose":
151            permutation = OnnxModel.get_node_attribute(parent, "perm")
152            if permutation == [0, 3, 1, 2]:
153                self.remove_if_safe(parent, input_name_to_nodes)
154                return parent
155        return None
156
157    def fuse(self, node, input_name_to_nodes, output_name_to_node):
158        # This fusion requires shape information, so skip it if shape is not available.
159        if self.shape_infer_helper is None:
160            return
161
162        # Before Fusion:
163        #     (a)  (b)
164        #       \  /
165        #       Add
166        #       /\
167        #   (c)   Transpose([0,2,3,1])
168        #            \
169        #          GroupNorm
170        #             |
171        #            (d)
172        #
173        # After Fusion:
174        #           (a)              (b)
175        #             \              /
176        #   Transpose([0,2,3,1])   Transpose([0,2,3,1])
177        #                \        /
178        #              SkipGroupNorm
179        #                  /    \
180        #                 /    Transpose([0, 3, 1, 2])
181        #                /        \
182        #               (d)       (c)
183        nodes = self.model.match_parent_path(node, ["Transpose", "Add"], [0, 0], output_name_to_node)
184        if nodes is None:
185            return
186
187        (transpose, add) = nodes
188        if transpose in self.nodes_to_remove or add in self.nodes_to_remove:
189            return
190
191        if self.has_multiple_consumers(transpose.output[0], input_name_to_nodes):
192            return
193
194        permutation = OnnxModel.get_node_attribute(transpose, "perm")
195        if permutation != [0, 2, 3, 1]:
196            return
197
198        inputs = []
199        bias = None
200        for i in range(2):
201            matched_transpose = self.match_transpose_from_nhwc(add.input[i], input_name_to_nodes, output_name_to_node)
202            if matched_transpose:
203                # When there is an Transpose node before Add (see examples in match_bias_path), we do not need to
204                # insert another Transpose node. The existing Transpose node will be removed in prune_graph if it
205                # has only one consumer.
206                inputs.append(matched_transpose.input[0])
207                # See whether it match bias pattern.
208                if bias is None:
209                    bias = self.match_bias_path(matched_transpose, input_name_to_nodes, output_name_to_node)
210            else:
211                # Otherwise, insert a Transpose node before Add.
212                new_transpose = self.create_transpose_node(add.input[i], [0, 2, 3, 1])
213                self.model.add_node(new_transpose, self.this_graph_name)
214                inputs.append(new_transpose.output[0])
215
216        skip, broadcast = self.get_skip_index(add, is_channel_last=False)
217        if skip < 0:
218            return
219
220        inputs = [inputs[1 - skip], node.input[1], node.input[2], inputs[skip]]
221        if bias:
222            inputs = [*inputs, bias]
223
224        outputs = node.output
225
226        new_node_name = self.model.create_node_name(self.fused_op_type, name_prefix="SkipGroupNorm")
227        if self.has_multiple_consumers(add.output[0], input_name_to_nodes):
228            add_out_name = new_node_name + "_add_out"
229            outputs.append(add_out_name)
230
231            # Insert a Transpose node after add output.
232            add_out_transpose = self.create_transpose_node(add_out_name, [0, 3, 1, 2], add.output[0])
233            self.model.add_node(add_out_transpose, self.this_graph_name)
234
235        skip_group_norm = helper.make_node(
236            self.fused_op_type,
237            inputs=inputs,
238            outputs=outputs,
239            name=new_node_name,
240        )
241        skip_group_norm.domain = "com.microsoft"
242
243        self.increase_counter(
244            f"SkipGroupNorm(add_out={int(len(outputs) > 1)} bias={int(bias is not None)} broadcast={int(broadcast)})"
245        )
246
247        # Pass attributes from GroupNorm node to SkipGroupNorm
248        for att in node.attribute:
249            skip_group_norm.attribute.extend([att])
250
251        self.nodes_to_remove.extend([add, transpose, node])
252        self.nodes_to_add.append(skip_group_norm)
253        self.node_name_to_graph_name[skip_group_norm.name] = self.this_graph_name
254        self.prune_graph = True
255 
codekingpro/portable-devtools · Team Ai