Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_constant_fold.py145 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_model import OnnxModel
11
12logger = getLogger(__name__)
13
14
15class FusionConstantFold(Fusion):
16    def __init__(self, model: OnnxModel):
17        super().__init__(model, "", ["Transpose"])
18        self.count = 0
19
20    def apply(self):
21        super().apply()
22        if self.count > 0:
23            logger.info(f"Constant Folded: {self.count}")
24
25    def fuse(self, node, input_name_to_nodes, output_name_to_node):
26        """
27        Apply multiple fusions on Transpose nodes that can be constant folded.
28        """
29        self.fuse_1(node, input_name_to_nodes, output_name_to_node)
30        self.fuse_2(node, input_name_to_nodes, output_name_to_node)
31
32    def fuse_1(self, node, input_name_to_nodes, output_name_to_node):
33        """
34        Constant fold any initializer data representing a MatMul's
35        weights that are stored in a Transpose op
36
37        Ex: Transpose --> Gemm or Transpose --> MatMul
38        """
39        # Check if Transpose node only has one input and one output
40        if len(node.input) != 1 or len(node.output) != 1:
41            logger.debug("fuse_constant_fold: node has more than one input or output")
42            return
43
44        # Check if input is initializer data
45        proto = self.model.get_initializer(node.input[0])
46        if proto is None:
47            logger.debug("fuse_constant_fold: failed to identify initializer input")
48            return
49
50        # Check that all nodes using input are Transpose ops that also only use the initializer data as input
51        skip = False
52        for child_node in input_name_to_nodes[node.input[0]]:
53            if not (child_node.op_type == "Transpose" and len(node.input) == 1):
54                skip = True
55                break
56        if skip:
57            logger.debug("fuse_constant_fold: other non-Transpose nodes use the initializer")
58            return
59
60        # Check that all nodes using output are Gemm or MatMul ops
61        for child_node in input_name_to_nodes[node.output[0]]:
62            if not (child_node.op_type == "Gemm" or child_node.op_type == "MatMul"):
63                skip = True
64                break
65        if skip:
66            logger.debug("fuse_constant_fold: other non-Gemm and non-MatMul nodes use the transposed data")
67            return
68
69        # Check if initializer data is 2D
70        weight = NumpyHelper.to_array(proto)
71        if len(weight.shape) != 2:
72            logger.debug("fuse_constant_fold: shape of initializer data is not 2D")
73            return
74
75        # Remove old TensorProto and add new TensorProto while re-using same name
76        name = proto.name
77        dtype = proto.data_type
78        self.remove_initializer(proto)
79        self.add_initializer(
80            name=name,
81            data_type=dtype,
82            dims=[weight.shape[1], weight.shape[0]],
83            vals=weight.T,
84        )
85
86        # Update weights input to be the initializer name and not
87        # the output of the Transpose op
88        for child_node in input_name_to_nodes[node.output[0]]:
89            for i in range(len(child_node.input)):
90                if child_node.input[i] == node.output[0]:
91                    child_node.input[i] = node.input[0]
92
93                    if child_node.op_type == "Gemm" and (i == 0 or i == 1):
94                        # Ensure that transA/transB is set to 0 in Gemm
95                        key = "transA" if i == 0 else "transB"
96                        for j, attr_key in enumerate(child_node.attribute):
97                            if attr_key.name == key:
98                                child_node.attribute[j].i = 0
99
100        # Add node to list of nodes to remove
101        self.nodes_to_remove.append(node)
102        self.count += 1
103
104    def fuse_2(self, node, input_name_to_nodes, output_name_to_node):
105        """
106        Constant fold any Transpose --> Transpose ops since the root input
107        is the final result
108
109        Ex: root_input --> Transpose --> Transpose --> next_node to root_input --> next_node
110        """
111        # Check if Transpose node only has one input and one output
112        if len(node.input) != 1 or len(node.output) != 1:
113            logger.debug("fuse_constant_fold: node has more than one input or output")
114            return
115
116        # Check if parent node is Transpose node with only one input and one output
117        parent_node = self.model.match_parent(node, "Transpose", 0)
118        if parent_node is None:
119            logger.debug("fuse_constant_fold: failed to identify parent Transpose node")
120            return
121        if len(parent_node.input) != 1 or len(parent_node.output) != 1:
122            logger.debug("fuse_constant_fold: parent node has more than one input or output")
123            return
124
125        node_perm = node.attribute[0].ints
126        parent_node_perm = parent_node.attribute[0].ints
127
128        if node_perm != parent_node_perm:
129            logger.debug("fuse_constant_fold: Transpose node permutations aren't identical")
130            return
131
132        # For nodes that use output of child Transpose node as an input,
133        # replace that input with root_input
134        root_input = parent_node.input[0]
135        output_nodes = input_name_to_nodes[node.output[0]]
136        for output_node in output_nodes:
137            for i, input_ in enumerate(output_node.input):
138                if input_ == node.output[0]:
139                    output_node.input[i] = root_input
140
141        # Add node to list of nodes to remove
142        self.nodes_to_remove.append(node)
143        self.nodes_to_remove.append(parent_node)
144        self.count += 1
145 
codekingpro/portable-devtools · Team Ai