Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_bias_add.py58 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 numpy import ndarray
9from onnx import helper
10from onnx_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionBiasAdd(Fusion):
16    def __init__(self, model: OnnxModel):
17        super().__init__(model, "BiasAdd", "Add")
18
19    def fuse(self, add_node, input_name_to_nodes: dict, output_name_to_node: dict):
20        """
21        Fuse Add bias and Add skip connection into BiasAdd
22        """
23
24        nodes = self.model.match_parent_path(
25            add_node,
26            ["Add", "MatMul", "BiasSplitGelu", "MatMul", "SkipLayerNormalization"],
27            [0, None, 0, 0, 0],
28            output_name_to_node,
29        )
30
31        if nodes is None:
32            return
33
34        bias_node = nodes[0]
35        skip_layer_norm = nodes[-1]
36
37        # Check skip connection is from SkipLayerNormalization output
38        if add_node.input[1] not in skip_layer_norm.output:
39            return
40
41        bias_index, bias_value = self.model.get_constant_input(bias_node)
42        if not (isinstance(bias_index, int) and (bias_value is not None) and isinstance(bias_value, ndarray)):
43            return
44        if bias_value.ndim != 1:
45            return
46
47        self.nodes_to_remove.extend([add_node, bias_node])
48        node_name = self.model.create_node_name("BiasAdd")
49        fused_node = helper.make_node(
50            "BiasAdd",
51            inputs=[bias_node.input[1 - bias_index], bias_node.input[bias_index], add_node.input[1]],
52            outputs=[add_node.output[0]],
53            name=node_name,
54        )
55        fused_node.domain = "com.microsoft"
56        self.nodes_to_add.append(fused_node)
57        self.node_name_to_graph_name[node_name] = self.this_graph_name
58 
codekingpro/portable-devtools · Team Ai