Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_skiplayernorm.py259 linesDownload Raw Back to transformers
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 import helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16def _is_broadcast_skip(input_shape, skip_shape):
17    """Check if skip_shape can broadcast to input_shape for SkipLayerNormalization.
18
19    The kernel supports: input 3D (B,S,H) with skip 3D (1,S,H) or skip 2D (S,H).
20    """
21    if len(input_shape) != 3:
22        return False
23    if len(skip_shape) == 3:
24        return skip_shape[0] == 1 and skip_shape[1] == input_shape[1] and skip_shape[2] == input_shape[2]
25    if len(skip_shape) == 2:
26        return skip_shape[0] == input_shape[1] and skip_shape[1] == input_shape[2]
27    return False
28
29
30class FusionSkipLayerNormalization(Fusion):
31    """
32    Fuse Add + LayerNormalization into one node: SkipLayerNormalization.
33    Supports broadcasting of the skip input: (1, sequence_length, hidden_size)
34    or (sequence_length, hidden_size) will be broadcast to match the input shape.
35    """
36
37    def __init__(
38        self,
39        model: OnnxModel,
40        fused_op_type: str = "SkipLayerNormalization",
41        search_op_types: str = "LayerNormalization",
42        shape_infer: bool = True,
43    ):
44        super().__init__(model, fused_op_type, search_op_types)
45        if shape_infer:
46            # Update shape inference is needed since other fusions might add new edge which does not have shape info yet.
47            self.shape_infer_helper = self.model.infer_runtime_shape({"batch_size": 4, "seq_len": 7}, update=True)
48            if self.shape_infer_helper is None:
49                # TODO(tianleiwu): support subgraph in shape inference.
50                logger.warning("symbolic shape inference disabled or failed.")
51
52    def get_skip_index(self, add):
53        """Identify which Add input is the skip tensor (the one that may broadcast).
54
55        Returns (skip_index, broadcast):
56            skip_index: 0 or 1 (which Add input is skip), -1 if incompatible
57            broadcast: True if broadcasting is needed
58        """
59        shape_a = self.shape_infer_helper.get_edge_shape(add.input[0])
60        shape_b = self.shape_infer_helper.get_edge_shape(add.input[1])
61        if shape_a is None or shape_b is None:
62            return -1, False
63
64        if shape_a == shape_b:
65            return (1, False) if len(shape_a) == 3 else (-1, False)
66
67        # Check if b is a broadcastable skip for a
68        if _is_broadcast_skip(shape_a, shape_b):
69            return 1, True
70        # Check if a is a broadcastable skip for b
71        if _is_broadcast_skip(shape_b, shape_a):
72            return 0, True
73
74        return -1, False
75
76    def fuse(self, node, input_name_to_nodes, output_name_to_node):
77        add = self.model.get_parent(node, 0, output_name_to_node)
78
79        # In some models there is input_ids->gather->add->LayerNorm and one of input of the
80        # add node is initializer with fixed shape which should not be fused into SkipLayerNorm
81        if add is None or add.op_type != "Add":
82            return
83
84        # The number of inputs of add should be 2
85        if len(add.input) != 2:
86            return
87
88        for add_input in add.input:
89            if self.model.get_initializer(add_input) is not None:
90                return
91
92        # To avoid an Add node have two children of LayerNormalization, we shall only fuse one SkipLayerNormalization
93        if add in self.nodes_to_remove:
94            return
95
96        # Root Mean Square Layer Normalization
97        simplified = node.op_type == "SimplifiedLayerNormalization"
98
99        skip_index = 1  # default: add.input[1] is the skip
100        _broadcast = False
101
102        if hasattr(self, "shape_infer_helper"):
103            if self.shape_infer_helper is not None:
104                skip_index, _broadcast = self.get_skip_index(add)
105                if skip_index < 0:
106                    logger.debug(
107                        "skip SkipLayerNormalization fusion since shapes of inputs (%s, %s) are not compatible",
108                        add.input[0],
109                        add.input[1],
110                    )
111                    return
112            else:
113                logger.debug("skip SkipLayerNormalization fusion since symbolic shape inference failed")
114                return
115
116        gather_path = self.model.match_parent_path(add, ["Gather"], [None])
117        if gather_path is not None and self.model.find_graph_input(gather_path[0].input[1]) is None:
118            if self.model.match_parent_path(gather_path[0], ["ConstantOfShape"], [1]) is None:
119                return
120
121        # When broadcasting is needed, check that neither Add input comes from a Gather
122        # (embedding lookup). Embedding Add+LayerNorm should be fused by EmbedLayerNormalization
123        # later in the pipeline, not as SkipLayerNormalization.
124        if _broadcast:
125            for i in range(2):
126                parent = self.model.get_parent(add, i, output_name_to_node)
127                if parent is not None and parent.op_type == "Gather":
128                    logger.debug(
129                        "skip SkipLayerNormalization broadcast fusion since Add input %d comes from Gather (embedding)",
130                        i,
131                    )
132                    return
133
134        # This means that the residual Add before the LayerNormalization produces an output
135        # that is consumed by some other nodes or graph output other than the LayerNormalization itself
136        # We can still go ahead with the SkipLayerNormalization fusion but we need to
137        # preserve the output of Add and that needs to be produced by SkipLayerNormalization.
138        add_has_graph_output = self.model.find_graph_output(add.output[0]) is not None
139        residual_add_has_multiple_consumers = (
140            add_has_graph_output or len(self.model.get_children(add, input_name_to_nodes)) > 1
141        )
142
143        outputs_to_keep = node.output
144
145        if residual_add_has_multiple_consumers:
146            outputs_to_keep.extend([add.output[0]])
147
148        outputs = [node.output[0]]
149
150        # Skip the other optional outputs of SkipLayerNormalization before adding the Add's output
151        if residual_add_has_multiple_consumers:
152            outputs.extend(["", "", add.output[0]])
153
154        if self.model.is_safe_to_fuse_nodes([add, node], outputs_to_keep, input_name_to_nodes, output_name_to_node):
155            self.nodes_to_remove.extend([add, node])
156
157            input_index = 1 - skip_index
158            inputs = (
159                [add.input[input_index], add.input[skip_index], node.input[1], node.input[2]]
160                if not simplified
161                else [add.input[input_index], add.input[skip_index], node.input[1]]
162            )
163            normalize_node = helper.make_node(
164                self.fused_op_type,
165                inputs=inputs,
166                outputs=outputs,
167                name=self.model.create_node_name(self.fused_op_type, name_prefix="SkipLayerNorm"),
168            )
169            normalize_node.domain = "com.microsoft"
170
171            # Pass attribute "epsilon" from layernorm node to SkipLayerNormalization
172            for att in node.attribute:
173                if att.name == "epsilon":
174                    normalize_node.attribute.extend([att])
175
176            # Set default epsilon if no epsilon exists from layernorm
177            if len(normalize_node.attribute) == 0:
178                normalize_node.attribute.extend([helper.make_attribute("epsilon", 1.0e-12)])
179
180            self.nodes_to_add.append(normalize_node)
181            self.node_name_to_graph_name[normalize_node.name] = self.this_graph_name
182
183
184class FusionBiasSkipLayerNormalization(Fusion):
185    def __init__(self, model: OnnxModel):
186        super().__init__(model, "SkipLayerNormalization", "SkipLayerNormalization", "add bias")
187
188    def fuse(self, node, input_name_to_nodes, output_name_to_node):
189        if len(node.input) != 4:
190            return
191
192        return_indice = []
193        nodes = self.model.match_parent_path(node, ["Add", "MatMul"], [None, None], output_name_to_node, return_indice)
194        if nodes is not None:
195            (add, _matmul) = nodes
196        else:
197            # In case of fp16, we could have a Cast between the MatMul and the bias Add
198            return_indice = []
199            nodes = self.model.match_parent_path(
200                node, ["Add", "Cast", "MatMul"], [None, None, None], output_name_to_node, return_indice
201            )
202            if nodes is not None:
203                (add, _cast, _matmul) = nodes
204            else:
205                return
206
207        assert len(return_indice) == 2 or len(return_indice) == 3
208        add_input_index = return_indice[0]
209        if add_input_index >= 2:
210            return
211        sln_input = add.input[return_indice[1]]
212        bias_input = add.input[1 - return_indice[1]]
213        skip_input = node.input[1 - add_input_index]
214
215        # bias should be one dimension
216        initializer = self.model.get_initializer(bias_input)
217        if initializer is None:
218            return
219        bias_weight = NumpyHelper.to_array(initializer)
220        if bias_weight is None:
221            logger.debug("Bias weight not found")
222            return
223        if len(bias_weight.shape) != 1:
224            logger.debug("Bias weight is not 1D")
225            return
226
227        subgraph_nodes = [node, add]
228        if not self.model.is_safe_to_fuse_nodes(subgraph_nodes, node.output, input_name_to_nodes, output_name_to_node):
229            logger.debug("Skip fusing SkipLayerNormalization with Bias since it is not safe")
230            return
231
232        self.nodes_to_remove.extend(subgraph_nodes)
233        inputs = [
234            sln_input,
235            skip_input,
236            node.input[2],
237            node.input[3],
238            bias_input,
239        ]
240        new_node = helper.make_node(
241            "SkipLayerNormalization",
242            inputs=inputs,
243            outputs=node.output,
244            name=self.model.create_node_name("SkipLayerNormalization", "SkipLayerNorm_AddBias_"),
245        )
246        new_node.domain = "com.microsoft"
247
248        # Pass attribute "epsilon" from skiplayernorm node to skiplayernorm(add bias)
249        for att in node.attribute:
250            if att.name == "epsilon":
251                new_node.attribute.extend([att])
252
253        # Set default epsilon if no epsilon exists from skiplayernorm
254        if len(new_node.attribute) == 0:
255            new_node.attribute.extend([helper.make_attribute("epsilon", 1.0e-12)])
256
257        self.nodes_to_add.append(new_node)
258        self.node_name_to_graph_name[new_node.name] = self.this_graph_name
259 
codekingpro/portable-devtools · Team Ai