Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
fusion_reshape.py174 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
8import numpy as np
9from fusion_base import Fusion
10from onnx import TensorProto, helper
11from onnx_model import OnnxModel
12
13logger = getLogger(__name__)
14
15
16class FusionReshape(Fusion):
17    def __init__(self, model: OnnxModel):
18        super().__init__(model, "Reshape", "Reshape")
19        self.prune_graph: bool = False
20
21    def replace_reshape_node(self, shape, reshape_node, concat_node):
22        shape_value = np.asarray(shape, dtype=np.int64)
23        constant_shape_name = self.model.create_node_name("Constant", "constant_shape")
24        new_node = helper.make_node(
25            "Constant",
26            inputs=[],
27            outputs=[constant_shape_name],
28            value=helper.make_tensor(
29                name="const_tensor",
30                data_type=TensorProto.INT64,
31                dims=shape_value.shape,
32                vals=bytes(shape_value),
33                raw=True,
34            ),
35        )
36        reshape_node.input[1] = constant_shape_name
37        reshape_node.name = self.model.create_node_name("Reshape", "Reshape_Fuse")
38        self.nodes_to_remove.extend([concat_node])
39        self.nodes_to_add.append(new_node)
40        self.node_name_to_graph_name[new_node.name] = self.this_graph_name
41
42    def fuse(self, reshape_node, input_name_to_nodes, output_name_to_node):
43        if reshape_node.input[1] not in output_name_to_node:
44            return
45
46        concat_node = output_name_to_node[reshape_node.input[1]]
47        if concat_node.op_type != "Concat" or len(concat_node.input) < 3 or len(concat_node.input) > 4:
48            return
49
50        path0 = self.model.match_parent_path(
51            concat_node,
52            ["Unsqueeze", "Gather", "Shape"],
53            [0, 0, 0],
54            output_name_to_node,
55        )
56        if path0 is None:
57            return
58
59        (unsqueeze_0, gather_0, shape_0) = path0
60
61        path1 = self.model.match_parent_path(
62            concat_node,
63            ["Unsqueeze", "Gather", "Shape"],
64            [1, 0, 0],
65            output_name_to_node,
66        )
67        if path1 is None:
68            return
69        (unsqueeze_1, gather_1, shape_1) = path1
70
71        shape = []
72        gather_value = self.model.get_constant_value(gather_0.input[1])
73        if gather_value == 0:
74            shape.append(0)
75
76        gather_value = self.model.get_constant_value(gather_1.input[1])
77        if gather_value == 1:
78            shape.append(0)
79
80        if len(shape) != 2:
81            return
82
83        path2 = []
84        path3 = []
85        shape_nodes = [shape_0, shape_1]
86        if len(concat_node.input) == 3 and self.model.get_constant_value(concat_node.input[2]) is None:
87            path2 = self.model.match_parent_path(
88                concat_node,
89                ["Unsqueeze", "Mul", "Gather", "Shape"],
90                [2, 0, 0, 0],
91                output_name_to_node,
92            )
93            if path2 is None:
94                path2 = self.model.match_parent_path(
95                    concat_node,
96                    ["Unsqueeze", "Mul", "Squeeze", "Slice", "Shape"],
97                    [2, 0, 0, 0, 0],
98                    output_name_to_node,
99                )  # GPT2 exported by PyTorch 1.4 with opset_version=11
100                if path2 is None:
101                    return
102
103            path3 = self.model.match_parent_path(
104                concat_node,
105                ["Unsqueeze", "Mul", "Gather", "Shape"],
106                [2, 0, 1, 0],
107                output_name_to_node,
108            )
109            if path3 is None:
110                path3 = self.model.match_parent_path(
111                    concat_node,
112                    ["Unsqueeze", "Mul", "Squeeze", "Slice", "Shape"],
113                    [2, 0, 1, 0, 0],
114                    output_name_to_node,
115                )  # GPT2 exported by PyTorch 1.4 with opset_version=11
116                if path3 is None:
117                    return
118
119            shape_nodes.extend([path2[-1], path3[-1]])
120            shape.append(-1)
121        elif len(concat_node.input) > 2:
122            concat_value = self.model.get_constant_value(concat_node.input[2])
123            if concat_value is None:
124                return
125            if isinstance(concat_value, np.ndarray):
126                shape.extend(concat_value.tolist())
127            else:
128                shape.append(concat_value)
129
130        if len(concat_node.input) == 4 and self.model.get_constant_value(concat_node.input[3]) is None:
131            if -1 in shape:
132                return
133
134            path2 = self.model.match_parent_path(
135                concat_node,
136                ["Unsqueeze", "Div", "Gather", "Shape"],
137                [3, 0, 0, 0],
138                output_name_to_node,
139            )
140            if path2 is None:
141                path2 = self.model.match_parent_path(
142                    concat_node,
143                    ["Unsqueeze", "Div", "Squeeze", "Slice", "Shape"],
144                    [3, 0, 0, 0, 0],
145                    output_name_to_node,
146                )  # GPT2 exported by PyTorch 1.4 with opset_version=11
147                if path2 is None:
148                    return
149            shape_nodes.extend([path2[-1]])
150            shape.append(-1)
151        elif len(concat_node.input) > 3:
152            concat_value = self.model.get_constant_value(concat_node.input[3])
153            if concat_value is None:
154                return
155
156            if isinstance(concat_value, np.ndarray):
157                shape.extend(concat_value.tolist())
158            else:
159                shape.append(concat_value)
160
161        root_input = reshape_node.input[0]
162        same_shape_input = True
163        for shape_node in shape_nodes:
164            if shape_node.input[0] != root_input:
165                same_shape_input = False
166
167        if not same_shape_input:
168            return
169
170        self.replace_reshape_node(shape, reshape_node, concat_node)
171
172        # TODO(tlwu): Subgraph blocks pruning un-used nodes. Add code to remove un-used nodes safely.
173        self.prune_graph = True
174 
codekingpro/portable-devtools · Team Ai