Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_model.py601 linesDownload Raw Back to quantization
1# --------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5from pathlib import Path
6
7import onnx
8import onnx.helper as onnx_helper
9import onnx.numpy_helper as onnx_numpy_helper
10from onnx.onnx_pb import ModelProto
11
12from .quant_utils import attribute_to_kwarg, find_by_name
13
14
15def _clean_initializers_helper(graph, model):
16    """Clean unused initializers from graph.
17
18    Returns:
19        A cleaned graph without unused initializers
20        A list of tensor names, which are not produced by this graph and its subgraphes
21    """
22    requesting_tensor_names = set()
23    requesting_tensor_names.update(input_name for node in graph.node for input_name in node.input if input_name)
24    requesting_tensor_names.update(g_out.name for g_out in graph.output if g_out.name)
25
26    new_nodes = []
27    for node in graph.node:
28        new_node = node
29        graph_attrs = [
30            attr
31            for attr in node.attribute
32            if attr.type == onnx.AttributeProto.GRAPH or attr.type == onnx.AttributeProto.GRAPHS
33        ]
34        if graph_attrs:
35            kwargs = {}
36            for attr in node.attribute:
37                new_attribute = {}
38                if attr.type == onnx.AttributeProto.GRAPH:
39                    (
40                        cleaned_sub_graph,
41                        sub_requesting_tensor_names,
42                    ) = _clean_initializers_helper(attr.g, model)
43                    new_attribute = {attr.name: cleaned_sub_graph}
44                    requesting_tensor_names.update(sub_requesting_tensor_names)
45                elif attr.type == onnx.AttributeProto.GRAPHS:
46                    cleaned_graphes = []
47                    for subgraph in attr.graphs:
48                        (
49                            cleaned_sub_graph,
50                            sub_requesting_tensor_names,
51                        ) = _clean_initializers_helper(subgraph, model)
52                        cleaned_graphes.append(cleaned_sub_graph)
53                        requesting_tensor_names.update(sub_requesting_tensor_names)
54                    new_attribute = {attr.name: cleaned_graphes}
55                else:
56                    new_attribute = attribute_to_kwarg(attr)
57                kwargs.update(new_attribute)
58            new_node = onnx_helper.make_node(node.op_type, node.input, node.output, name=node.name, **kwargs)
59        new_nodes.append(new_node)
60
61    graph.ClearField("node")
62    graph.node.extend(new_nodes)
63
64    requesting_tensor_names.difference_update(output for node in graph.node for output in node.output)
65
66    unused_initializer = []
67    for initializer in graph.initializer:
68        if initializer.name in requesting_tensor_names:
69            requesting_tensor_names.remove(initializer.name)
70        else:
71            # mark it to remove, remove here directly will cause mis-behavier
72            unused_initializer.append(initializer)
73
74    name_to_input = {input.name: input for input in graph.input}
75    for initializer in unused_initializer:
76        graph.initializer.remove(initializer)
77        if initializer.name in name_to_input:
78            try:
79                graph.input.remove(name_to_input[initializer.name])
80            except StopIteration:
81                if model.ir_version < 4:
82                    print(f"Warning: invalid weight name {initializer.name} found in the graph (not a graph input)")
83
84    requesting_tensor_names.difference_update(input.name for input in graph.input)
85
86    return graph, requesting_tensor_names
87
88
89class ONNXModel:
90    def __init__(self, model: ModelProto):
91        self.model = model
92
93    def nodes(self):
94        return self.model.graph.node
95
96    def initializer(self):
97        return self.model.graph.initializer
98
99    def initializer_extend(self, inits):
100        if len(inits) == 0:
101            raise ValueError("Can add an empty list.")
102        for init in self.initializer():
103            self._check_init(init, "gain")
104        for init in inits:
105            self._check_init(init)
106            self.model.graph.initializer.append(init)
107
108    def graph(self):
109        return self.model.graph
110
111    def ir_version(self):
112        return self.model.ir_version
113
114    def opset_import(self):
115        return self.model.opset_import
116
117    def set_opset_import(self, domain, version):
118        for opset in self.model.opset_import:
119            if opset.domain == domain:
120                opset.version = version
121                return
122
123        self.model.opset_import.extend([onnx_helper.make_opsetid(domain, version)])
124
125    def remove_node(self, node):
126        if node in self.model.graph.node:
127            self.model.graph.node.remove(node)
128
129    def remove_nodes(self, nodes_to_remove):
130        for node in nodes_to_remove:
131            self.remove_node(node)
132
133    def add_node(self, node):
134        self.model.graph.node.extend([self._check_node(node)])
135
136    def add_nodes(self, nodes_to_add):
137        for node in nodes_to_add:
138            self.add_node(node)
139
140    def add_initializer(self, tensor):
141        if find_by_name(tensor.name, self.model.graph.initializer) is None:
142            self._check_init(tensor)
143            self.model.graph.initializer.extend([tensor])
144
145    def get_initializer(self, name):
146        for tensor in self.model.graph.initializer:
147            if tensor.name == name:
148                return tensor
149        return None
150
151    def find_graph_input(self, input_name):
152        for input in self.model.graph.input:
153            if input.name == input_name:
154                return input
155        return None
156
157    def find_graph_output(self, output_name):
158        for output in self.model.graph.output:
159            if output.name == output_name:
160                return output
161        return None
162
163    def get_tensor_type(self, tensor_name: str):
164        tensor_type_map = {obj.name: obj.type for obj in self.model.graph.value_info}
165
166        if tensor_name in tensor_type_map:
167            return tensor_type_map[tensor_name].tensor_type
168
169        g_input = self.find_graph_input(tensor_name)
170        if g_input:
171            return g_input.type.tensor_type
172
173        g_output = self.find_graph_output(tensor_name)
174        if g_output:
175            return g_output.type.tensor_type
176
177        return None
178
179    def get_constant_value(self, output_name):
180        for node in self.model.graph.node:
181            if node.op_type == "Constant":
182                if node.output[0] == output_name:
183                    for attr in node.attribute:
184                        if attr.name == "value":
185                            return onnx_numpy_helper.to_array(attr.t)
186
187        # Fallback to initializer since constant folding may have been applied.
188        initializer = self.get_initializer(output_name)
189        if initializer is not None:
190            return onnx_numpy_helper.to_array(initializer)
191
192        return None
193
194    def get_initializer_name_set(self):
195        return {initializer.name for initializer in self.model.graph.initializer}
196
197    def remove_initializer(self, tensor):
198        if tensor in self.model.graph.initializer:
199            self.model.graph.initializer.remove(tensor)
200            for input in self.model.graph.input:
201                if input.name == tensor.name:
202                    self.model.graph.input.remove(input)
203                    break
204
205    def remove_initializers(self, init_to_remove):
206        for initializer in init_to_remove:
207            self.remove_initializer(initializer)
208
209    def get_non_initializer_inputs(self):
210        initializer_names = self.get_initializer_name_set()
211        non_initializer_inputs = set()
212        for input in self.model.graph.input:
213            if input.name not in initializer_names:
214                non_initializer_inputs.add(input.name)
215        return non_initializer_inputs
216
217    def input_name_to_nodes(self):
218        input_name_to_nodes = {}
219        for node in self.model.graph.node:
220            for input_name in node.input:
221                if input_name:  # Could be empty when it is optional
222                    if input_name not in input_name_to_nodes:
223                        input_name_to_nodes[input_name] = [node]
224                    else:
225                        input_name_to_nodes[input_name].append(node)
226        return input_name_to_nodes
227
228    def output_name_to_node(self):
229        output_name_to_node = {}
230        for node in self.model.graph.node:
231            for output_name in node.output:
232                if output_name:  # Could be empty when it is optional
233                    output_name_to_node[output_name] = node
234        return output_name_to_node
235
236    def get_children(self, node, input_name_to_nodes=None):
237        if input_name_to_nodes is None:
238            input_name_to_nodes = self.input_name_to_nodes()
239
240        children = []
241        for output in node.output:
242            if output in input_name_to_nodes:
243                for node in input_name_to_nodes[output]:
244                    children.append(node)  # noqa: PERF402
245        return children
246
247    def get_parents(self, node, output_name_to_node=None):
248        if output_name_to_node is None:
249            output_name_to_node = self.output_name_to_node()
250
251        parents = []
252        for input in node.input:
253            if input in output_name_to_node:
254                parents.append(output_name_to_node[input])
255        return parents
256
257    def get_parent(self, node, idx, output_name_to_node=None):
258        if output_name_to_node is None:
259            output_name_to_node = self.output_name_to_node()
260
261        if len(node.input) <= idx:
262            return None
263
264        input = node.input[idx]
265        if input not in output_name_to_node:
266            return None
267
268        return output_name_to_node[input]
269
270    def find_node_by_name(self, node_name, new_nodes_list, graph):
271        """Find out if a node exists in a graph or a node is in the
272        new set of nodes created during quantization.
273
274        Returns:
275            The node found or None.
276        """
277        graph_nodes_list = list(graph.node)  # deep copy
278        graph_nodes_list.extend(new_nodes_list)
279        node = find_by_name(node_name, graph_nodes_list)
280        return node
281
282    def get_largest_node_name_suffix(self, node_name_prefix):
283        """
284        Gets the largest node name (int) suffix for all node names that begin with `node_name_prefix`.
285        Example: for nodes my_prefix_0 and my_prefix_3, this method returns 3.
286        """
287        suffix = -1
288
289        for node in self.model.graph.node:
290            if node.name and node.name.startswith(node_name_prefix):
291                try:
292                    index = int(node.name[len(node_name_prefix) :])
293                    suffix = max(index, suffix)
294                except ValueError:
295                    continue
296
297        return suffix
298
299    def get_largest_initializer_name_suffix(self, initializer_name_prefix):
300        """
301        Gets the largest initializer name integer suffix for all initializer names that begin
302        with `initializer_name_prefix`. This can be used to create unique initializer names.
303
304        Example: for initializer names 'my_weight_0' and 'my_weight_3', this method returns 3 if
305                 `initializer_name_prefix` is 'my_weight_'.
306        """
307        suffix = -1
308
309        for initializer in self.model.graph.initializer:
310            if initializer.name.startswith(initializer_name_prefix):
311                try:
312                    index = int(initializer.name[len(initializer_name_prefix) :])
313                    suffix = max(index, suffix)
314                except ValueError:
315                    continue
316
317        return suffix
318
319    def find_nodes_by_initializer(self, graph, initializer):
320        """
321        Find all nodes with given initializer as an input.
322        """
323        nodes = []
324        for node in graph.node:
325            for node_input in node.input:
326                if node_input == initializer.name:
327                    nodes.append(node)
328        return nodes
329
330    @staticmethod
331    def __get_initializer(name, graph_path):
332        for gid in range(len(graph_path) - 1, -1, -1):
333            graph = graph_path[gid]
334            for tensor in graph.initializer:
335                if tensor.name == name:
336                    return tensor, graph
337        return None, None
338
339    @staticmethod
340    def __replace_gemm_with_matmul(graph_path):
341        new_nodes = []
342        graph = graph_path[-1]
343        for node in graph.node:
344            graph_attrs = [attr for attr in node.attribute if attr.type == 5 or attr.type == 10]
345            if graph_attrs:
346                kwargs = {}
347                for attr in node.attribute:
348                    if attr.type == 5:
349                        graph_path.append(attr.g)
350                        kv = {attr.name: ONNXModel.__replace_gemm_with_matmul(graph_path)}
351                    elif attr.type == 10:
352                        value = []
353                        for subgraph in attr.graphs:
354                            graph_path.append(subgraph)
355                            value.extend([ONNXModel.__replace_gemm_with_matmul(graph_path)])
356                        kv = {attr.name: value}
357                    else:
358                        kv = attribute_to_kwarg(attr)
359                    kwargs.update(kv)
360                node = onnx_helper.make_node(  # noqa: PLW2901
361                    node.op_type, node.input, node.output, name=node.name, **kwargs
362                )
363
364            if node.op_type == "Gemm":
365                alpha = 1.0
366                beta = 1.0
367                transA = 0  # noqa: N806
368                transB = 0  # noqa: N806
369                for attr in node.attribute:
370                    if attr.name == "alpha":
371                        alpha = onnx_helper.get_attribute_value(attr)
372                    elif attr.name == "beta":
373                        beta = onnx_helper.get_attribute_value(attr)
374                    elif attr.name == "transA":
375                        transA = onnx_helper.get_attribute_value(attr)  # noqa: N806
376                    elif attr.name == "transB":
377                        transB = onnx_helper.get_attribute_value(attr)  # noqa: N806
378                if alpha == 1.0 and beta == 1.0 and transA == 0:
379                    inputB = node.input[1]  # noqa: N806
380                    if transB == 1:
381                        B, Bs_graph = ONNXModel.__get_initializer(node.input[1], graph_path)  # noqa: N806
382                        if B:
383                            # assume B is not used by any other node
384                            B_array = onnx_numpy_helper.to_array(B)  # noqa: N806
385                            B_trans = onnx_numpy_helper.from_array(B_array.T)  # noqa: N806
386                            B_trans.name = B.name
387                            Bs_graph.initializer.remove(B)
388                            for input in Bs_graph.input:
389                                if input.name == inputB:
390                                    Bs_graph.input.remove(input)
391                                    break
392                            Bs_graph.initializer.extend([B_trans])
393                        else:
394                            inputB += "_Transposed"  # noqa: N806
395                            transpose_node = onnx_helper.make_node(
396                                "Transpose",
397                                inputs=[node.input[1]],
398                                outputs=[inputB],
399                                name=node.name + "_Transpose" if node.name else "",
400                            )
401                            new_nodes.append(transpose_node)
402
403                    matmul_node = onnx_helper.make_node(
404                        "MatMul",
405                        inputs=[node.input[0], inputB],
406                        outputs=[node.output[0] + ("_MatMul" if len(node.input) > 2 else "")],
407                        name=node.name + "_MatMul" if node.name else "",
408                    )
409                    new_nodes.append(matmul_node)
410
411                    if len(node.input) > 2:
412                        add_node = onnx_helper.make_node(
413                            "Add",
414                            inputs=[node.output[0] + "_MatMul", node.input[2]],
415                            outputs=node.output,
416                            name=node.name + "_Add" if node.name else "",
417                        )
418                        new_nodes.append(add_node)
419
420                # unsupported
421                else:
422                    new_nodes.append(node)
423
424            # not GEMM
425            else:
426                new_nodes.append(node)
427
428        graph.ClearField("node")
429        graph.node.extend(new_nodes)
430        graph_path.pop()
431        return graph
432
433    def replace_gemm_with_matmul(self):
434        graph_path = [self.graph()]
435        ONNXModel.__replace_gemm_with_matmul(graph_path)
436
437    def save_model_to_file(self, output_path, use_external_data_format=False):
438        """
439        Save model to external data, which is needed for model size > 2GB
440        """
441        self.topological_sort()
442        if use_external_data_format:
443            onnx.external_data_helper.convert_model_to_external_data(
444                self.model,
445                all_tensors_to_one_file=True,
446                location=Path(output_path).name + ".data",
447                convert_attribute=True,
448            )
449        for init in self.model.graph.initializer:
450            self._check_init(init, "end")
451        onnx.save_model(self.model, output_path)
452
453    @staticmethod
454    def replace_node_input(node, old_input_name, new_input_name):
455        assert isinstance(old_input_name, str) and isinstance(new_input_name, str)
456        for j in range(len(node.input)):
457            if node.input[j] == old_input_name:
458                node.input[j] = new_input_name
459
460    def replace_input_of_all_nodes(self, old_input_name, new_input_name):
461        for node in self.model.graph.node:
462            ONNXModel.replace_node_input(node, old_input_name, new_input_name)
463
464    def replace_input_of_nodes(self, old_input_name, new_input_name, node_names_set):
465        for node in self.model.graph.node:
466            if node.name in node_names_set:
467                ONNXModel.replace_node_input(node, old_input_name, new_input_name)
468
469    @staticmethod
470    def replace_node_output(node, old_output_name, new_output_name):
471        assert isinstance(old_output_name, str) and isinstance(new_output_name, str)
472        for j in range(len(node.output)):
473            if node.output[j] == old_output_name:
474                node.output[j] = new_output_name
475
476    def replace_output_of_all_nodes(self, old_output_name, new_output_name):
477        for node in self.model.graph.node:
478            ONNXModel.replace_node_output(node, old_output_name, new_output_name)
479
480    def replace_output_of_nodes(self, old_output_name, new_output_name, node_names_set):
481        for node in self.model.graph.node:
482            if node.name in node_names_set:
483                ONNXModel.replace_node_output(node, old_output_name, new_output_name)
484
485    def remove_unused_constant(self):
486        input_name_to_nodes = self.input_name_to_nodes()
487
488        # remove unused constant
489        unused_nodes = []
490        nodes = self.nodes()
491        for node in nodes:
492            if (
493                node.op_type == "Constant"
494                and not self.is_graph_output(node.output[0])
495                and node.output[0] not in input_name_to_nodes
496            ):
497                unused_nodes.append(node)
498
499        self.remove_nodes(unused_nodes)
500
501        ununsed_weights = []
502        for w in self.initializer():
503            if w.name not in input_name_to_nodes and not self.is_graph_output(w.name):
504                ununsed_weights.append(w)
505                # Remove from graph.input
506                for graph_input in self.graph().input:
507                    if graph_input.name == w.name:
508                        self.graph().input.remove(graph_input)
509
510        self.remove_initializers(ununsed_weights)
511
512    def is_graph_output(self, output_name):
513        return any(output.name == output_name for output in self.model.graph.output)
514
515    def is_graph_input(self, tensor_name: str) -> bool:
516        return any(input.name == tensor_name for input in self.model.graph.input)
517
518    # TODO:use OnnxModel.graph_topological_sort(self.model.graph) from transformers.onnx_model
519    # Currently it breaks Openvino/Linux training gpu pipeline so hold off for 1.8 release
520    def topological_sort(self):
521        deps_count = [0] * len(self.nodes())  # dependency count of each node
522        deps_to_nodes = {}  # input to node indice
523        sorted_nodes = []  # initialize sorted_nodes
524        for node_idx, node in enumerate(self.nodes()):
525            # CANNOT use len(node.input) directly because input can be optional
526            deps_count[node_idx] = sum(1 for _ in node.input if _)
527            if deps_count[node_idx] == 0:  # Constant doesn't depend on any inputs
528                sorted_nodes.append(self.nodes()[node_idx])
529                continue
530
531            for input_name in node.input:
532                if not input_name:
533                    continue
534                if input_name not in deps_to_nodes:
535                    deps_to_nodes[input_name] = [node_idx]
536                else:
537                    deps_to_nodes[input_name].append(node_idx)
538
539        initializer_names = [init.name for init in self.initializer()]
540        graph_input_names = [input.name for input in self.model.graph.input]
541        input_names = initializer_names + graph_input_names
542        input_names.sort()
543        prev_input_name = None
544        for input_name in input_names:
545            if prev_input_name == input_name:
546                continue
547
548            prev_input_name = input_name
549            if input_name in deps_to_nodes:
550                for node_idx in deps_to_nodes[input_name]:
551                    deps_count[node_idx] = deps_count[node_idx] - 1
552                    if deps_count[node_idx] == 0:
553                        sorted_nodes.append(self.nodes()[node_idx])
554
555        start = 0
556        end = len(sorted_nodes)
557
558        while start < end:
559            for output in sorted_nodes[start].output:
560                if output in deps_to_nodes:
561                    for node_idx in deps_to_nodes[output]:
562                        deps_count[node_idx] = deps_count[node_idx] - 1
563                        if deps_count[node_idx] == 0:
564                            sorted_nodes.append(self.nodes()[node_idx])
565                            end = end + 1
566            start = start + 1
567
568        assert end == len(self.graph().node), "Graph is not a DAG"
569        self.graph().ClearField("node")
570        self.graph().node.extend(sorted_nodes)
571
572    def clean_initializers(self):
573        return _clean_initializers_helper(self.graph(), self.model)
574
575    def _check_init(self, init, test=None):
576        if init.data_type == onnx.TensorProto.FLOAT8E4M3FN:
577            if init.HasField("raw_data"):
578                b = list(init.raw_data)
579                if any((i & 127) == 127 for i in b):
580                    raise ValueError(f"Initializer {init.name!r} has nan.")
581        return init
582
583    def _check_node(self, node):
584        """
585        A quantization to float 8 does not use quantized bias but float 16 bias.
586        This function checks that DequantizeLinear is not used to
587        dequantize from float 16.
588        """
589        if node.op_type == "DequantizeLinear":
590            zero_point = node.input[2]
591            init = self.get_initializer(zero_point)
592            dtype = init.data_type
593            if dtype in {
594                onnx.TensorProto.FLOAT16,
595                onnx.TensorProto.FLOAT,
596                onnx.TensorProto.DOUBLE,
597                onnx.TensorProto.BFLOAT16,
598            }:
599                raise RuntimeError(f"Unsupported DequantizeLinear operator, dequantization from {dtype}.")
600        return node
601 
codekingpro/portable-devtools · Team Ai