Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
dynamo_onnx_helper.py206 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation. All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5from collections.abc import Sequence
6from logging import getLogger
7from typing import Any
8
9import numpy as np
10import onnx
11from onnx import helper
12from onnx_model import OnnxModel
13
14logger = getLogger(__name__)
15
16
17class DynamoOnnxHelper:
18    """
19    Helper class for processing ONNX models exported by Torch Dynamo.
20    """
21
22    def __init__(self, model: onnx.ModelProto):
23        self.model = OnnxModel(model)
24
25    def update_edges(self, edge_mapping: dict) -> None:
26        """
27        Updates the edges in the model according to the given mapping.
28        """
29        for node in self.model.model.graph.node:
30            for i in range(len(node.input)):
31                if node.input[i] in edge_mapping:
32                    node.input[i] = edge_mapping[node.input[i]]
33            for i in range(len(node.output)):
34                if node.output[i] in edge_mapping:
35                    node.output[i] = edge_mapping[node.output[i]]
36
37        for graph_input in self.model.model.graph.input:
38            if graph_input.name in edge_mapping:
39                graph_input.name = edge_mapping[graph_input.name]
40        for graph_output in self.model.model.graph.output:
41            if graph_output.name in edge_mapping:
42                graph_output.name = edge_mapping[graph_output.name]
43
44    def unroll_function(self, func_name: str) -> None:
45        """
46        Unrolls the function with the given name in the model.
47        """
48        logger.debug(f"Unrolling function {func_name}...")
49        nodes_to_remove = []
50        nodes_to_add = []
51        edges_to_remove = []
52        edges_to_add = []
53        for node in self.model.model.graph.node:
54            if node.op_type == func_name:
55                nodes_to_remove.append(node)
56                edges_to_remove.extend(list(node.input) + list(node.output))
57
58        func_to_remove = None
59        for f in self.model.model.functions:
60            if f.name == func_name:
61                nodes_to_add.extend(list(f.node))
62                edges_to_add.extend(list(f.input) + list(f.output))
63                func_to_remove = f
64
65        assert len(edges_to_remove) == len(edges_to_add)
66
67        for node in nodes_to_remove:
68            self.model.model.graph.node.remove(node)
69        for node in nodes_to_add:
70            self.model.model.graph.node.append(node)
71        if func_to_remove is not None:
72            self.model.model.functions.remove(func_to_remove)
73
74        edge_mapping = {}
75        for i in range(len(edges_to_remove)):
76            k = edges_to_remove[i]
77            v = edges_to_add[i]
78            if k != v:
79                edge_mapping[k] = v
80
81        return self.update_edges(edge_mapping)
82
83    def remove_function(self, func_name: str, input_id: int, output_id: int) -> None:
84        """
85        Removes the function in the model.
86        """
87        edge_mapping = {}
88        nodes_to_remove = []
89        for node in self.model.model.graph.node:
90            if node.op_type.find(func_name) != -1:
91                edge_mapping[node.input[input_id]] = node.output[output_id]
92                nodes_to_remove.append(node)
93        for node in nodes_to_remove:
94            self.model.model.graph.node.remove(node)
95
96        self.update_edges(edge_mapping)
97
98    def remove_dropout_layer(self) -> None:
99        """
100        Removes the dropout layer in the model.
101        """
102        logger.debug("Removing dropout layer...")
103        self.remove_function("Dropout", 0, 0)
104
105    def remove_lm_head_layer(self) -> None:
106        """
107        Removes the LM head layer in the model.
108        """
109        logger.debug("Removing LM head layer...")
110        # bugbug: need to copy the right vi over
111        self.remove_function("Linear_lm_head", 2, 0)
112
113    def add_initializer(self, name: str, data_type: int, dims: Sequence[int], vals: Any, raw: bool = True):
114        if raw:
115            np_type = helper.tensor_dtype_to_np_dtype(data_type)
116            if not isinstance(vals, np.ndarray):
117                bytes = np.array(vals, dtype=np_type).tobytes()
118            else:
119                bytes = vals.astype(np_type).tobytes()
120            tensor = helper.make_tensor(
121                name=name,
122                data_type=data_type,
123                dims=dims,
124                vals=bytes,
125                raw=True,
126            )
127        else:
128            tensor = helper.make_tensor(
129                name=name,
130                data_type=data_type,
131                dims=dims,
132                vals=vals,
133                raw=False,
134            )
135
136        self.model.add_initializer(tensor)
137        return tensor
138
139    def convert_constants_to_initializers(self, min_size: int = 1) -> None:
140        """
141        Converts Constant ops of size [min_size] or higher to initializers
142        """
143        logger.debug(f"Converting constants greater than size {min_size} to initializers")
144
145        constant_nodes = self.model.get_nodes_by_op_type("Constant")
146        nodes_to_remove = []
147
148        for node in constant_nodes:
149            # Get info from Constant op
150            np_data = self.model.get_constant_value(node.output[0])
151
152            # Skip if there are less than [min_size] elements
153            if np_data is None or np_data.size < min_size:
154                continue
155
156            # Add new initializer with same name as Constant op's output
157            for att in node.attribute:
158                if att.name == "value":
159                    self.add_initializer(
160                        name=node.output[0],
161                        data_type=att.t.data_type,
162                        dims=list(np_data.shape),
163                        vals=np_data,
164                    )
165                    break
166
167            nodes_to_remove.append(node)
168
169        # Remove Constant ops from graph
170        self.model.remove_nodes(nodes_to_remove)
171
172    def clear_metadata(self) -> None:
173        """
174        Clear metadata fields in all nodes
175        """
176        for graph in self.model.graphs():
177            graph.ClearField("metadata_props")
178        for node in self.model.nodes():
179            node.ClearField("metadata_props")
180
181    @staticmethod
182    def fold_transpose_initializers(model) -> None:
183        """
184        Constant fold Transpose initializers without changing the initializer names
185        """
186        from onnxscript import ir  # noqa: PLC0415
187
188        for name, initializer in model.graph.initializers.items():
189            user_nodes = initializer.consumers()
190            if len(user_nodes) == 1 and user_nodes[0].op_type == "Transpose":
191                transpose_node = user_nodes[0]
192                perm = transpose_node.attributes.get("perm")
193                if perm is None:
194                    transposed_tensor = ir.tensor(initializer.const_value.numpy().transpose())
195                else:
196                    transposed_tensor = ir.tensor(initializer.const_value.numpy().transpose(perm.as_ints()))
197                new_initializer = ir.Value(
198                    name=initializer.name,
199                    shape=transposed_tensor.shape,
200                    type=ir.TensorType(transposed_tensor.dtype),
201                    const_value=transposed_tensor,
202                )
203                ir.convenience.replace_all_uses_with(transpose_node.outputs[0], new_initializer)
204                model.graph.initializers[name] = new_initializer
205                transpose_node.graph.remove(transpose_node, safe=True)
206 
codekingpro/portable-devtools · Team Ai