Team Ai
Datasetpublic

codekingpro/portable-devtools

sourceHugging Faceupdated 5mo agoView on Hugging Face
1likes14kdownloads
onnx_model.py1637 linesDownload Raw Back to transformers
1# -------------------------------------------------------------------------
2# Copyright (c) Microsoft Corporation.  All rights reserved.
3# Licensed under the MIT License.
4# --------------------------------------------------------------------------
5
6import itertools
7import logging
8import os
9import sys
10from collections import deque
11from pathlib import Path
12
13from float16 import convert_float_to_float16
14from onnx import (
15    AttributeProto,
16    GraphProto,
17    ModelProto,
18    NodeProto,
19    TensorProto,
20    ValueInfoProto,
21    helper,
22    numpy_helper,
23    save_model,
24)
25from onnx.external_data_helper import load_external_data_for_tensor, uses_external_data
26from shape_infer_helper import SymbolicShapeInferenceHelper
27
28logger = logging.getLogger(__name__)
29
30
31class OnnxModel:
32    def __init__(self, model):
33        self.initialize(model)
34
35    def initialize(self, model):
36        self.model: ModelProto = model
37        self._node_name_suffix: dict[str, int] = {}  # key is node name prefix, value is the last suffix generated
38        self.shape_infer_helper: SymbolicShapeInferenceHelper = None
39        self.enable_shape_infer: bool = True
40        self.all_graphs: list[GraphProto] | None = None
41
42        # Cache of shape and data type from onnx graph to speed up optimization.
43        # Be careful that fusion shall not reuse node output name for different shape/type (in adding/removing nodes)
44        # Note that these do not cache the symbolic shape inference result.
45        self._dtype_dict: dict[str, int] | None = None
46        self._shape_dict: dict[str, list] | None = None
47
48    def disable_shape_inference(self):
49        self.enable_shape_infer = False
50
51    def infer_runtime_shape(self, dynamic_axis_mapping={}, update=False):  # noqa: B006
52        if self.enable_shape_infer:
53            if self.shape_infer_helper is None or update:
54                self.shape_infer_helper = SymbolicShapeInferenceHelper(self.model)
55
56            try:
57                if self.shape_infer_helper.infer(dynamic_axis_mapping):
58                    return self.shape_infer_helper
59            except Exception:
60                self.enable_shape_infer = False  # disable shape inference to suppress same error message.
61                print("failed in shape inference", sys.exc_info()[0])
62
63        return None
64
65    def input_name_to_nodes(self, exclude_subgraphs=False):
66        input_name_to_nodes = {}
67        nodes_to_search = self.nodes() if not exclude_subgraphs else self.model.graph.node
68        for node in nodes_to_search:
69            for input_name in node.input:
70                if input_name:  # could be empty when it is optional
71                    if input_name not in input_name_to_nodes:
72                        input_name_to_nodes[input_name] = [node]
73                    else:
74                        input_name_to_nodes[input_name].append(node)
75        return input_name_to_nodes
76
77    def output_name_to_node(self, exclude_subgraphs=False):
78        output_name_to_node = {}
79        nodes_to_search = self.nodes() if not exclude_subgraphs else self.model.graph.node
80        for node in nodes_to_search:
81            for output_name in node.output:
82                if output_name:  # could be empty when it is optional
83                    output_name_to_node[output_name] = node
84        return output_name_to_node
85
86    def functions(self):
87        all_functions = [list(self.model.functions)]
88        return all_functions
89
90    def nodes(self):
91        all_nodes = []
92        for graph in self.graphs():
93            for node in graph.node:
94                all_nodes.append(node)  # noqa: PERF402
95        return all_nodes
96
97    def graph(self):
98        return self.model.graph
99
100    def graphs(self):
101        if self.all_graphs is not None:
102            return self.all_graphs
103        self.all_graphs = []
104        graph_queue = [self.model.graph]
105        while graph_queue:
106            graph = graph_queue.pop(0)
107            self.all_graphs.append(graph)
108            for node in graph.node:
109                for attr in node.attribute:
110                    if attr.type == AttributeProto.AttributeType.GRAPH:
111                        assert isinstance(attr.g, GraphProto)
112                        graph_queue.append(attr.g)
113                    if attr.type == AttributeProto.AttributeType.GRAPHS:
114                        for g in attr.graphs:
115                            assert isinstance(g, GraphProto)
116                            graph_queue.append(g)
117        return self.all_graphs
118
119    def get_graphs_input_names(self):
120        input_names = []
121        for graph in self.graphs():
122            for input in graph.input:
123                input_names.append(input.name)
124        return input_names
125
126    def get_graphs_output_names(self):
127        output_names = []
128        for graph in self.graphs():
129            for output in graph.output:
130                output_names.append(output.name)
131        return output_names
132
133    def get_graph_by_node(self, node):
134        for graph in self.graphs():
135            if node in graph.node:
136                return graph
137        return None
138
139    def get_graph_by_name(self, graph_name):
140        for graph in self.graphs():
141            if graph_name == graph.name:
142                return graph
143        return None
144
145    def get_topological_insert_id(self, graph, outputs):
146        for idx, node in enumerate(graph.node):
147            for input in node.input:
148                if input in outputs:
149                    return idx
150        return len(graph.node)
151
152    def remove_node(self, node):
153        for graph in self.graphs():
154            if node in graph.node:
155                graph.node.remove(node)
156                return
157        logger.warning("Failed to remove node %s", node)  # It might be a bug to hit this line.
158
159    def remove_nodes(self, nodes_to_remove):
160        for node in nodes_to_remove:
161            self.remove_node(node)
162
163    def add_node(self, node, graph_name=None):
164        if graph_name is None or graph_name == self.model.graph.name:
165            self.model.graph.node.extend([node])
166        else:
167            graph = self.get_graph_by_name(graph_name)
168            insert_idx = self.get_topological_insert_id(graph, node.output)
169            graph.node.insert(insert_idx, node)
170
171    def add_nodes(self, nodes_to_add, node_name_to_graph_name=None):
172        if node_name_to_graph_name is None:
173            self.model.graph.node.extend(nodes_to_add)
174        else:
175            for node in nodes_to_add:
176                graph_name = node_name_to_graph_name[node.name]
177                self.add_node(node, graph_name)
178
179    def add_initializer(self, tensor, graph_name=None):
180        if graph_name is None or graph_name == self.model.graph.name:
181            self.model.graph.initializer.extend([tensor])
182        else:
183            graph = self.get_graph_by_name(graph_name)
184            graph.initializer.extend([tensor])
185
186    def add_input(self, input, graph_name=None):
187        if graph_name is None or graph_name == self.model.graph.name:
188            self.model.graph.input.extend([input])
189        else:
190            graph = self.get_graph_by_name(graph_name)
191            graph.input.extend([input])
192
193    @staticmethod
194    def replace_node_input(node, old_input_name, new_input_name):
195        assert isinstance(old_input_name, str) and isinstance(new_input_name, str)
196        for j in range(len(node.input)):
197            if node.input[j] == old_input_name:
198                node.input[j] = new_input_name
199
200    def replace_input_of_all_nodes(self, old_input_name, new_input_name):
201        for node in self.nodes():
202            OnnxModel.replace_node_input(node, old_input_name, new_input_name)
203
204    @staticmethod
205    def replace_node_output(node, old_output_name, new_output_name):
206        assert isinstance(old_output_name, str) and isinstance(new_output_name, str)
207        for j in range(len(node.output)):
208            if node.output[j] == old_output_name:
209                node.output[j] = new_output_name
210
211    def replace_output_of_all_nodes(self, old_output_name, new_output_name):
212        # This function shall be used carefully. For example:
213        #       Add --[old_name]--> Cast ---> [new_name]
214        #        |
215        #        +----[old_name]--> Transpose -->
216        # If we want to remove the Cast node: replace output of Add to new_name is not enough;
217        # The input of Transpose shall also be updated to new_name.
218        for node in self.model.graph.node:
219            OnnxModel.replace_node_output(node, old_output_name, new_output_name)
220
221    def get_initializer(self, name):
222        for graph in self.graphs():
223            for tensor in graph.initializer:
224                if tensor.name == name:
225                    return tensor
226        return None
227
228    def get_nodes_by_op_type(self, op_type):
229        nodes = []
230        for node in self.nodes():
231            if node.op_type == op_type:
232                nodes.append(node)
233        return nodes
234
235    def get_children(self, node, input_name_to_nodes=None, output_index=None):
236        if input_name_to_nodes is None:
237            input_name_to_nodes = self.input_name_to_nodes()
238
239        children = []
240        if output_index is not None:
241            if output_index < len(node.output):
242                output = node.output[output_index]
243                if output in input_name_to_nodes:
244                    children = list(input_name_to_nodes[output])
245        else:
246            for output in node.output:
247                if output in input_name_to_nodes:
248                    children.extend(input_name_to_nodes[output])
249
250        return children
251
252    def get_parents(self, node, output_name_to_node=None):
253        if output_name_to_node is None:
254            output_name_to_node = self.output_name_to_node()
255
256        parents = []
257        for input in node.input:
258            if input in output_name_to_node:
259                parents.append(output_name_to_node[input])
260        return parents
261
262    def get_parent(self, node, i, output_name_to_node=None):
263        if output_name_to_node is None:
264            output_name_to_node = self.output_name_to_node()
265
266        if len(node.input) <= i:
267            return None
268
269        input = node.input[i]
270        if input not in output_name_to_node:
271            return None
272
273        return output_name_to_node[input]
274
275    def match_first_parent(self, node, parent_op_type, output_name_to_node, exclude=[]):  # noqa: B006
276        """
277        Find parent node based on constraints on op_type.
278
279        Args:
280            node (str): current node name.
281            parent_op_type (str): constraint of parent node op_type.
282            output_name_to_node (dict): dictionary with output name as key, and node as value.
283            exclude (list): list of nodes that are excluded (not allowed to match as parent).
284
285        Returns:
286            parent: The matched parent node. None if not found.
287            index: The input index of matched parent node. None if not found.
288        """
289        for i, input in enumerate(node.input):
290            if input in output_name_to_node:
291                parent = output_name_to_node[input]
292                if parent.op_type == parent_op_type and parent not in exclude:
293                    return parent, i
294                else:
295                    logger.debug(f"To find first {parent_op_type}, current {parent.op_type}")
296        return None, None
297
298    def match_parent(
299        self,
300        node,
301        parent_op_type,
302        input_index=None,
303        output_name_to_node=None,
304        exclude=[],  # noqa: B006
305        return_indice=None,
306    ):
307        """
308        Find parent node based on constraints on op_type and index.
309        When input_index is None, we will find the first parent node based on constraints,
310        and return_indice will be appended the corresponding input index.
311
312        Args:
313            node (str): current node name.
314            parent_op_type (str): constraint of parent node op_type.
315            input_index (int or None): only check the parent given input index of current node.
316            output_name_to_node (dict): dictionary with output name as key, and node as value.
317            exclude (list): list of nodes that are excluded (not allowed to match as parent).
318            return_indice (list): a list to append the input index when input_index is None.
319
320        Returns:
321            parent: The matched parent node.
322        """
323        assert node is not None
324        assert input_index is None or input_index >= 0
325
326        if output_name_to_node is None:
327            output_name_to_node = self.output_name_to_node()
328
329        if input_index is None:
330            parent, index = self.match_first_parent(node, parent_op_type, output_name_to_node, exclude)
331            if return_indice is not None:
332                return_indice.append(index)
333            return parent
334
335        if input_index >= len(node.input):
336            logger.debug(f"input_index {input_index} >= node inputs {len(node.input)}")
337            return None
338
339        parent = self.get_parent(node, input_index, output_name_to_node)
340        if parent is not None and parent.op_type == parent_op_type and parent not in exclude:
341            return parent
342
343        if parent is not None:
344            logger.debug(f"Expect {parent_op_type}, Got {parent.op_type}")
345
346        return None
347
348    def match_parent_paths(self, node, paths, output_name_to_node):
349        for i, path in enumerate(paths):
350            assert isinstance(path, (list, tuple))
351            return_indice = []
352            matched = self.match_parent_path(node, path[0], path[1], output_name_to_node, return_indice)
353            if matched:
354                return i, matched, return_indice
355        return -1, None, None
356
357    def match_parent_paths_all(self, node, paths, output_name_to_node):
358        match_i, matches, return_indices = [], [], []
359        for i, path in enumerate(paths):
360            assert isinstance(path, (list, tuple))
361            return_indice = []
362            matched = self.match_parent_path(node, path[0], path[1], output_name_to_node, return_indice)
363            if matched:
364                match_i.append(i)
365                matches.append(matched)
366                return_indices.append(return_indice)
367        return match_i, matches, return_indices
368
369    def match_parent_path(
370        self,
371        node,
372        parent_op_types,
373        parent_input_index=None,
374        output_name_to_node=None,
375        return_indice=None,
376    ):
377        """
378        Find a sequence of input edges based on constraints on parent op_type and index.
379        When input_index is None, we will find the first parent node based on constraints,
380        and return_indice will be appended the corresponding input index.
381
382        Args:
383            node (str): current node name.
384            parent_op_types (str): constraint of parent node op_type of each input edge.
385            parent_input_index (list): constraint of input index of each input edge. None means no constraint.
386            output_name_to_node (dict): dictionary with output name as key, and node as value.
387            return_indice (list): a list to append the input index
388                                  When there is no constraint on input index of an edge.
389
390        Returns:
391            parents: a list of matched parent node.
392        """
393        if parent_input_index is not None:
394            assert len(parent_input_index) == len(parent_op_types)
395
396        if output_name_to_node is None:
397            output_name_to_node = self.output_name_to_node()
398
399        current_node = node
400        matched_parents = []
401        for i, op_type in enumerate(parent_op_types):
402            matched_parent = self.match_parent(
403                current_node,
404                op_type,
405                parent_input_index[i] if parent_input_index is not None else None,
406                output_name_to_node,
407                exclude=[],
408                return_indice=return_indice,
409            )
410            if matched_parent is None:
411                if parent_input_index is not None:
412                    logger.debug(
413                        f"Failed to match index={i} parent_input_index={parent_input_index[i]} op_type={op_type}",
414                        stack_info=True,
415                    )
416                else:
417                    logger.debug(f"Failed to match index={i} op_type={op_type}", stack_info=True)
418                return None
419
420            matched_parents.append(matched_parent)
421            current_node = matched_parent
422
423        return matched_parents
424
425    def find_first_child_by_type(self, node, child_type, input_name_to_nodes=None, recursive=True):
426        children = self.get_children(node, input_name_to_nodes)
427        dq = deque(children)
428        while len(dq) > 0:
429            current_node = dq.pop()
430            if current_node.op_type == child_type:
431                return current_node
432
433            if recursive:
434                children = self.get_children(current_node, input_name_to_nodes)
435                for child in children:
436                    dq.appendleft(child)
437
438        return None
439
440    def match_child_path(
441        self,
442        node,
443        child_op_types,
444        edges: list[tuple[int, int]] | None = None,
445        input_name_to_nodes=None,
446        exclude=[],  # noqa: B006
447    ):
448        """
449        Find a sequence of input edges based on constraints on parent op_type and index.
450        Note that we use greedy approach and only consider the first matched child, so it has chance to miss matching.
451
452        Args:
453            node (str): current node name.
454            child_op_types (str): constraint of child node op_type of each input edge.
455            edges (list): each edge is represented by two integers: output index of parent node, input index of child node.
456                         None means no constraint.
457            exclude(list): list of nodes that are excluded (not allowed to match as child).
458
459        Returns:
460            children: a list of matched children node.
461        """
462        if edges is not None:
463            assert len(edges) == len(child_op_types)
464            for edge in edges:
465                assert (
466                    isinstance(edge, tuple) and len(edge) == 2 and isinstance(edge[0], int) and isinstance(edge[1], int)
467                )
468
469        if input_name_to_nodes is None:
470            input_name_to_nodes = self.input_name_to_nodes()
471
472        current_node = node
473        matched_children = []
474        for i, op_type in enumerate(child_op_types):
475            matched_child = None
476
477            if edges is None:
478                children_nodes = self.get_children(current_node, input_name_to_nodes=input_name_to_nodes)
479            else:
480                children_nodes = self.get_children(
481                    current_node, input_name_to_nodes=input_name_to_nodes, output_index=edges[i][0]
482                )
483
484            for child in children_nodes:
485                if child.op_type == op_type and child not in exclude:
486                    if edges is not None and child.input[edges[i][1]] != current_node.output[edges[i][0]]:
487                        continue
488
489                    # Here we use greedy approach and only consider the first matched child.
490                    # TODO: match recursively if we encounter cases that the correct child is not the first matched.
491                    matched_child = child
492                    break
493
494            if matched_child is None:
495                logger.debug(f"Failed to match child {i} op_type={op_type}", stack_info=True)
496                return None
497
498            matched_children.append(matched_child)
499            current_node = matched_child
500
501        return matched_children
502
503    def find_first_parent_by_type(self, node, parent_type, output_name_to_node=None, recursive=True):
504        if output_name_to_node is None:
505            output_name_to_node = self.output_name_to_node()
506
507        parents = self.get_parents(node, output_name_to_node)
508        dq = deque(parents)
509        while len(dq) > 0:
510            current_node = dq.pop()
511            if current_node.op_type == parent_type:
512                return current_node
513
514            if recursive:
515                parents = self.get_parents(current_node, output_name_to_node)
516                for parent in parents:
517                    dq.appendleft(parent)
518
519        return None
520
521    def get_constant_value(self, output_name):
522        for node in self.get_nodes_by_op_type("Constant"):
523            if node.output[0] == output_name:
524                for att in node.attribute:
525                    if att.name == "value":
526                        return numpy_helper.to_array(att.t)
527
528        # Fall back to intializer since constant folding might have been applied.
529        initializer = self.get_initializer(output_name)
530        if initializer is not None:
531            return numpy_helper.to_array(initializer)
532
533        return None
534
535    def get_constant_input(self, node):
536        for i, input in enumerate(node.input):
537            value = self.get_constant_value(input)
538            if value is not None:
539                return i, value
540
541        return None, None
542
543    def find_constant_input(self, node, expected_value, delta=0.000001):
544        i, value = self.get_constant_input(node)
545        if value is not None and value.size == 1 and abs(value - expected_value) < delta:
546            return i
547
548        return -1
549
550    def is_constant_with_specified_dimension(self, output_name, dimensions, description):
551        value = self.get_constant_value(output_name)
552        if value is None:
553            logger.debug(f"{description} {output_name} is not initializer.")
554            return False
555
556        if len(value.shape) != dimensions:
557            logger.debug(f"{description} {output_name} shall have {dimensions} dimensions. Got shape {value.shape}")
558            return False
559
560        return True
561
562    def has_constant_input(self, node, expected_value, delta=0.000001):
563        return self.find_constant_input(node, expected_value, delta) >= 0
564
565    def get_children_subgraph_nodes(self, root_node, stop_nodes, input_name_to_nodes=None):
566        if input_name_to_nodes is None:
567            input_name_to_nodes = self.input_name_to_nodes()
568
569        children = input_name_to_nodes[root_node.output[0]]
570
571        unique_nodes = []
572
573        dq = deque(children)
574        while len(dq) > 0:
575            current_node = dq.pop()
576            if current_node in stop_nodes:
577                continue
578
579            if current_node not in unique_nodes:
580                unique_nodes.append(current_node)
581
582                for output in current_node.output:
583                    if output in input_name_to_nodes:
584                        children = input_name_to_nodes[output]
585                        for child in children:
586                            dq.appendleft(child)
587
588        return unique_nodes
589
590    def tensor_shape_to_list(self, tensor_type):
591        """Convert tensor shape to list"""
592        shape_list = []
593        for d in tensor_type.shape.dim:
594            if d.HasField("dim_value"):
595                shape_list.append(d.dim_value)  # known dimension
596            elif d.HasField("dim_param"):
597                shape_list.append(d.dim_param)  # unknown dimension with symbolic name
598            else:
599                shape_list.append("?")  # shall not happen
600        return shape_list
601
602    def get_dtype(self, name: str, symbolic_shape_helper: SymbolicShapeInferenceHelper | None = None):
603        """Try get data type given a name (could be initializer, input or output of graph or node)."""
604
605        if self._dtype_dict is None:
606            self._dtype_dict = {}
607            for value_info in itertools.chain(
608                self.model.graph.value_info,
609                self.model.graph.input,
610                self.model.graph.output,
611            ):
612                self._dtype_dict[value_info.name] = value_info.type.tensor_type.elem_type
613
614            for initializer in self.model.graph.initializer:
615                if initializer.name not in self._dtype_dict:
616                    self._dtype_dict[initializer.name] = initializer.data_type
617
618        if name in self._dtype_dict:
619            return self._dtype_dict[name]
620
621        if symbolic_shape_helper is not None and name in symbolic_shape_helper.known_vi_:
622            value_info = symbolic_shape_helper.known_vi_[name]
623            return value_info.type.tensor_type.elem_type
624
625        return None
626
627    def get_shape(self, name: str, symbolic_shape_helper: SymbolicShapeInferenceHelper | None = None):
628        """Try get shape given a name (could be initializer, input or output of graph or node)."""
629
630        if self._shape_dict is None:
631            self._shape_dict = {}
632            for value_info in itertools.chain(
633                self.model.graph.value_info,
634                self.model.graph.input,
635                self.model.graph.output,
636            ):
637                if value_info.type.tensor_type.HasField("shape"):
638                    shape = []
639                    for dim in value_info.type.tensor_type.shape.dim:
640                        if dim.dim_param:
641                            shape.append(dim.dim_param)
642                        else:
643                            shape.append(dim.dim_value)
644                    self._shape_dict[value_info.name] = shape
645
646            for initializer in self.model.graph.initializer:
647                if initializer.name not in self._shape_dict:
648                    self._shape_dict[initializer.name] = initializer.dims
649
650        if name in self._shape_dict:
651            return self._shape_dict[name]
652
653        if symbolic_shape_helper is not None and name in symbolic_shape_helper.known_vi_:
654            value_info = symbolic_shape_helper.known_vi_[name]
655            return value_info.type.tensor_type.elem_type
656
657        return None
658
659    @staticmethod
660    def get_node_attribute(node: NodeProto, attribute_name: str):
661        for attr in node.attribute:
662            if attr.name == attribute_name:
663                value = helper.get_attribute_value(attr)
664                return value
665        return None
666
667    def remove_cascaded_cast_nodes(self):
668        """Remove Cast node that are followed by another Cast node like  --> Cast --> Cast -->
669        Note that this shall be used carefully since it might introduce semantic change.
670        For example, float -> int -> float could get different value than the original float value.
671        So, it is recommended to used only in post-processing of mixed precision conversion.
672        """
673        output_name_to_node = self.output_name_to_node()
674        removed_count = 0
675        for node in self.nodes():
676            if node.op_type == "Cast":
677                parent = self.get_parent(node, 0, output_name_to_node=output_name_to_node)
678                if parent and parent.op_type == "Cast":
679                    node.input[0] = parent.input[0]
680                    removed_count += 1
681
682        if removed_count > 0:
683            logger.info("Removed %d cascaded Cast nodes", removed_count)
684            self.prune_graph()
685
686    def remove_useless_cast_nodes(self):
687        """Remove cast nodes that are not needed: input and output has same data type."""
688        shape_infer = self.infer_runtime_shape(update=True)
689        if self.enable_shape_infer and shape_infer is None:
690            logger.warning("shape inference failed which might impact useless cast node detection.")
691
692        nodes_to_remove = []
693        for node in self.nodes():
694            if node.op_type == "Cast":
695                input_dtype = self.get_dtype(node.input[0], shape_infer)
696                output_dtype = self.get_dtype(node.output[0], shape_infer)
697                if input_dtype and input_dtype == output_dtype:
698                    nodes_to_remove.append(node)
699
700        if nodes_to_remove:
701            graph_input_names = set(self.get_graphs_input_names())
702            graph_output_names = set(self.get_graphs_output_names())
703            for node in nodes_to_remove:
704                if bool(set(node.output) & graph_output_names):
705                    if (not bool(set(node.input) & graph_input_names)) and len(
706                        self.input_name_to_nodes()[node.input[0]]
707                    ) == 1:
708                        self.replace_output_of_all_nodes(node.input[0], node.output[0])
709                    else:
710                        continue
711                else:
712                    self.replace_input_of_all_nodes(node.output[0], node.input[0])
713                self.remove_node(node)
714
715            logger.info(
716                "Removed %d Cast nodes with output type same as input",
717                len(nodes_to_remove),
718            )
719
720    def convert_model_float32_to_float16(self, cast_input_output=True):
721        logger.warning(
722            "The function convert_model_float32_to_float16 is deprecated. Use convert_float_to_float16 instead!"
723        )
724        self.convert_float_to_float16(use_symbolic_shape_infer=True, keep_io_types=cast_input_output)
725
726    def convert_float_to_float16(self, use_symbolic_shape_infer=True, **kwargs):
727        """Convert a model to half (default) or mixed precision.
728           To use mixed precision, user need specify which graph inputs, outputs, operator type
729           or list of nodes shall keep in float32.
730
731           Note that the conversion might not proceed without type information for the whole graph.
732
733           By default, we use symbolic shape inference to get type information. The benefit of symbolic shape inference
734           is that it could handle fused operators in com.microsoft domain. Those operators cannot be handled in onnx shape
735           inference so symbolic shape inference is recommended for optimized model.
736
737           When symbolic shape inference is used (even if it failed), ONNX shape inference will be disabled.
738
739           Note that onnx shape inference will fail for model larger than 2GB. For large model, you have to enable
740           symbolic shape inference. If your model is not optimized, you can also use model path to call
741           convert_float_to_float16 in float16.py (see https://github.com/microsoft/onnxruntime/pull/15067) to
742           avoid the 2GB limit.
743
744        Args:
745            use_symbolic_shape_infer (bool, optional): use symbolic shape inference instead of onnx shape inference.
746                                                       Defaults to True.
747            keep_io_types (Union[bool, List[str]], optional): boolean or a list of float32 input/output names.
748                                                              If True, model inputs/outputs should be left as float32.
749                                                              Defaults to True.
750            op_block_list (List[str], optional): List of operator types to leave as float32.
751                                                 Defaults to None, which will use `float16.DEFAULT_OP_BLOCK_LIST`.
752            node_block_list (List[str], optional): List of node names to leave as float32. Defaults to None.
753            force_fp16_initializers(bool): force converting all float initializers to float16.
754                                           Default to false.
755            min_positive_val (float, optional): minimal positive value. Defaults to 1e-7.
756            max_finite_val (float, optional): maximal finite value. Defaults to 1e4.
757            force_fp16_inputs(Dict[str, List[int]]): Force the conversion of the inputs of some operators to float16, even if
758                                                     this script's preference it to keep them in float32.
759        """
760        if "keep_io_types" not in kwargs:
761            kwargs["keep_io_types"] = True
762
763        model = self.model
764        if use_symbolic_shape_infer:
765            # Use symbolic shape inference since custom operators (like Gelu, SkipLayerNormalization etc)
766            # are not recognized by onnx shape inference.
767            shape_infer_helper = SymbolicShapeInferenceHelper(model)
768            try:
769                model_with_shape = shape_infer_helper.infer_shapes(model, auto_merge=True, guess_output_rank=False)
770
771                # auto_merge might cause issue (see https://github.com/microsoft/onnxruntime/issues/15521)
772                # we only merge tensor data type but not shape information back to the original onnx model.
773                # Note that float16 conversion need data type but not shape information.
774                if model_with_shape is not None:
775                    name_vi = {}
776                    for vi in model_with_shape.graph.value_info:
777                        if (
778                            hasattr(vi.type, "tensor_type")
779                            and hasattr(vi.type.tensor_type, "elem_type")
780                            and vi.type.tensor_type.elem_type != TensorProto.UNDEFINED
781                            and vi.name
782                        ):
783                            vi_copy = ValueInfoProto()
784                            vi_copy.CopyFrom(vi)
785                            if hasattr(vi_copy.type.tensor_type, "shape"):
786                                vi_copy.type.tensor_type.ClearField("shape")
787                            name_vi[vi.name] = vi_copy
788                    for vi in model.graph.value_info:
789                        if vi.name in name_vi:
790                            del name_vi[vi.name]
791                    for vi in name_vi.values():
792                        model.graph.value_info.append(vi)
793            except Exception:
794                logger.warning(
795                    "Failed to run symbolic shape inference. Please file an issue in https://github.com/microsoft/onnxruntime."
796                )
797
798        parameters = {"disable_shape_infer": use_symbolic_shape_infer}
799        parameters.update(
800            {
801                key: kwargs[key]
802                for key in [
803                    "keep_io_types",
804                    "min_positive_val",
805                    "max_finite_val",
806                    "op_block_list",
807                    "node_block_list",
808                    "force_fp16_initializers",
809                    "force_fp16_inputs",
810                    "use_bfloat16_as_blocked_nodes_dtype",
811                ]
812                if key in kwargs
813            }
814        )
815
816        fp16_model = convert_float_to_float16(model, **parameters)
817        self.initialize(fp16_model)
818
819        self.remove_cascaded_cast_nodes()
820
821        self.remove_useless_cast_nodes()
822
823    def create_node_name(self, op_type, name_prefix=None):
824        """Create a unique node name that starts with a prefix (default is operator type).
825           The name will not be duplicated with any name that generated or existed in current graphs.
826        Args:
827            op_type (str): operator type
828            name_prefix (str, optional): prefix of node name. Defaults to None.
829
830        Returns:
831            str: node name
832        """
833
834        if name_prefix:
835            prefix = name_prefix if name_prefix.endswith("_") else (name_prefix + "_")
836        else:
837            prefix = op_type + "_"
838
839        suffix: int = 0
840        if prefix in self._node_name_suffix:
841            suffix = self._node_name_suffix[prefix] + 1
842        else:
843            # Check existed node name only once for a prefix
844            # as we assume create_node_name is called for every new node in fusion.
845            for node in self.nodes():
846                if node.name and node.name.startswith(prefix):
847                    try:
848                        index = int(node.name[len(prefix) :])
849                        suffix = max(index + 1, suffix)
850                    except ValueError:
851                        continue
852
853        # Record the generated suffix so that we can avoid generating duplicated name.
854        self._node_name_suffix[prefix] = suffix
855
856        return prefix + str(suffix)
857
858    def find_graph_input(self, input_name):
859        for input in self.model.graph.input:
860            if input.name == input_name:
861                return input
862        return None
863
864    def find_graph_output(self, output_name):
865        for output in self.model.graph.output:
866            if output.name == output_name:
867                return output
868        return None
869
870    def get_parent_subgraph_nodes(self, node, stop_nodes, output_name_to_node=None):
871        if output_name_to_node is None:
872            output_name_to_node = self.output_name_to_node()
873
874        unique_nodes = []
875
876        parents = self.get_parents(node, output_name_to_node)
877        dq = deque(parents)
878        while len(dq) > 0:
879            current_node = dq.pop()
880            if current_node in stop_nodes:
881                continue
882
883            if current_node not in unique_nodes:
884                unique_nodes.append(current_node)
885
886                for input in current_node.input:
887                    if input in output_name_to_node:
888                        dq.appendleft(output_name_to_node[input])
889
890        return unique_nodes
891
892    def get_graph_inputs(self, current_node, recursive=False):
893        """
894        Find graph inputs that linked to current node.
895        """
896        graph_inputs = []
897        for input in current_node.input:
898            if self.find_graph_input(input) and input not in graph_inputs:
899                graph_inputs.append(input)
900
901        if recursive:
902            parent_nodes = self.get_parent_subgraph_nodes(current_node, [])
903            for node in parent_nodes:
904                for input in node.input:
905                    if self.find_graph_input(input) and input not in graph_inputs:
906                        graph_inputs.append(input)
907        return graph_inputs
908
909    @staticmethod
910    def input_index(node_output, child_node):
911        for index, input in enumerate(child_node.input):
912            if input == node_output:
913                return index
914        return -1
915
916    def remove_unused_constant(self):
917        input_name_to_nodes = self.input_name_to_nodes()
918
919        # remove unused constant
920        unused_nodes = []
921        nodes = self.nodes()
922        for node in nodes:
923            if node.op_type == "Constant" and node.output[0] not in input_name_to_nodes:
924                unused_nodes.append(node)
925
926        self.remove_nodes(unused_nodes)
927
928        if len(unused_nodes) > 0:
929            logger.debug(f"Removed unused constant nodes: {len(unused_nodes)}")
930
931    def _get_subgraph_inputs_of_node(self, node):
932        """
933        Get inputs to all nodes in all subgraphs of a node
934        """
935        # Note: This function only handles one-level subgraphs of child nodes.
936        subgraph_nodes_inputs = set()
937        for attr in node.attribute:
938            if attr.type == AttributeProto.GRAPH:
939                child_nodes = attr.g.node
940                for child_node in child_nodes:
941                    subgraph_nodes_inputs.update(child_node.input)
942        return subgraph_nodes_inputs
943
944    def _get_subgraph_nodes_and_inputs(self, ops_with_graph_attrs):
945        """
946        Get input names to all nodes in all subgraphs where subgraphs are
947        graph attributes of a node in the main graph
948        """
949        subgraph_nodes = list(filter(lambda node: node.op_type in ops_with_graph_attrs, self.model.graph.node))
950        subgraph_nodes_inputs = set()
951        for parent_node in subgraph_nodes:
952            subgraph_inputs_of_parent_node = self._get_subgraph_inputs_of_node(parent_node)
953            subgraph_nodes_inputs.update(subgraph_inputs_of_parent_node)
954        return subgraph_nodes, subgraph_nodes_inputs
955
956    def prune_graph(self, outputs=None, allow_remove_graph_inputs=True):
957        """
958        Prune graph to keep only required outputs. It removes unnecessary nodes that are not linked
959        (directly or indirectly) to any required output.
960
961        There is also an option to remove graph inputs that are not used to generate any required output.
962
963        Args:
964            outputs (list): a list of graph outputs to retain. If it is None, all graph outputs will be kept.
965            allow_remove_graph_inputs (bool): allow remove graph inputs.
966        """
967
968        keep_outputs = [output.name for output in self.model.graph.output] if outputs is None else outputs
969
970        input_name_to_nodes_for_main_graph = self.input_name_to_nodes(exclude_subgraphs=True)
971        output_name_to_node = self.output_name_to_node()
972
973        def get_first_output(node):
974            if node.output[0]:
975                return node.output[0]
976            return next(iter([o for o in node.output if o]), None)
977
978        if len(self.graphs()) > 1:
979            # Get input names for all nodes in all subgraphs
980            subgraph_nodes, subgraph_nodes_inputs = self._get_subgraph_nodes_and_inputs(
981                ops_with_graph_attrs={"Loop", "Scan", "If"}
982            )
983            if len(subgraph_nodes) == 0:
984                # TODO: support other ops such as `BeamSearch` that have subgraphs as op attributes
985                logger.debug("Skip prune_graph since graph has subgraph")
986                return
987
988            # For graphs with subgraphs, add dangling outputs from parent graph nodes to list of outputs to keep
989            for node in self.model.graph.node:
990                # TODO: This for-loop logic currently assumes that Loop/Scan/If nodes will not be
991                # pruned because their subgraphs are needed for computations. This might not be
992                # true in all cases.
993                if node in subgraph_nodes:
994                    continue
995
996                # Check if node output is an input of a subgraph node and not an input to a node in the main graph
997                for output in node.output:
998                    if output in subgraph_nodes_inputs and output not in input_name_to_nodes_for_main_graph:
999                        keep_outputs += [output]
1000
1001        # Keep track of nodes to keep. The key is first output of node, and the value is the node.
1002        output_to_node = {}
1003
1004        # Start from graph outputs, and find parent nodes recursively, and add nodes to the output_to_node dictionary.
1005        dq = deque()
1006        for output in keep_outputs:
1007            if output in output_name_to_node:
1008                dq.append(output_name_to_node[output])
1009        while len(dq) > 0:
1010            node = dq.pop()
1011            first_output = get_first_output(node)
1012            if first_output and (first_output not in output_to_node):
1013                output_to_node[first_output] = node
1014                for name in node.input:
1015                    if len(name) > 0 and (name in output_name_to_node) and (name not in output_to_node):
1016                        dq.appendleft(output_name_to_node[name])
1017
1018        # Keep only those nodes in the output_to_node dictionary.
1019        nodes_to_keep = []
1020        num_nodes_removed = 0
1021        for node in self.model.graph.node:
1022            first_output = get_first_output(node)
1023            kept_node = output_to_node.get(first_output)
1024
1025            # Need to double check the node since fused node might reuse output name of some nodes to be removed.
1026            # It is slow to compare whole node, so we compare op_type first to avoid comparing node in most cases.
1027            if kept_node and kept_node.op_type == node.op_type and kept_node == node:
1028                nodes_to_keep.append(node)
1029            else:
1030                num_nodes_removed += 1
1031
1032        self.all_graphs = (
1033            None  # to prevent pass-by-copy after ClearField(), forces the use of pass-by-reference instead
1034        )
1035        self.model.graph.ClearField("node")
1036        self.model.graph.node.extend(nodes_to_keep)
1037
1038        # Remove graph outputs not in list
1039        output_to_remove = []
1040        if outputs is not None:
1041            for output in self.model.graph.output:
1042                if output.name not in outputs:
1043                    output_to_remove.append(output)
1044            for output in output_to_remove:
1045                self.model.graph.output.remove(output)
1046
1047        # Remove graph inputs not used by any node.
1048        input_to_remove = []
1049        if allow_remove_graph_inputs:
1050            input_name_to_nodes = self.input_name_to_nodes()
1051            input_to_remove = [input for input in self.model.graph.input if input.name not in input_name_to_nodes]
1052            for name in input_to_remove:
1053                self.model.graph.input.remove(name)
1054
1055        if input_to_remove or output_to_remove or num_nodes_removed > 0:
1056            removed = []
1057            if input_to_remove:
1058                removed.append(f"{len(input_to_remove)} inputs")
1059            if output_to_remove:
1060                removed.append(f"{len(output_to_remove)} outputs")
1061            if num_nodes_removed > 0:
1062                removed.append(f"{num_nodes_removed} nodes")
1063            logger.info("Removed %s", ", ".join(removed))
1064
1065        self.update_graph()
1066
1067    def update_graph(self, verbose=False, allow_remove_graph_inputs=False):
1068        graph = self.model.graph
1069
1070        remaining_input_names = set()
1071        for node in graph.node:
1072            if node.op_type in ["Loop", "Scan", "If"]:
1073                # Add input names of nodes in subgraphs
1074                subgraph_inputs_of_node = self._get_subgraph_inputs_of_node(node)
1075                remaining_input_names.update(subgraph_inputs_of_node)
1076
1077            if node.op_type != "Constant":
1078                remaining_input_names.update(node.input)
1079        if verbose:
1080            logger.debug(f"remaining input names: {remaining_input_names}")
1081
1082        # remove graph input that is not used
1083        inputs_to_remove = []
1084        if allow_remove_graph_inputs:
1085            for input in graph.input:
1086                if input.name not in remaining_input_names:
1087                    inputs_to_remove.append(input)
1088            for input in inputs_to_remove:
1089                graph.input.remove(input)
1090
1091        names_to_remove = [input.name for input in inputs_to_remove]
1092        logger.debug(f"remove {len(inputs_to_remove)} unused inputs: {names_to_remove}")
1093
1094        # remove weights that are not used
1095        weights_to_remove = []
1096        weights_to_keep = []
1097        for initializer in graph.initializer:
1098            if initializer.name not in remaining_input_names and not self.find_graph_output(initializer.name):
1099                weights_to_remove.append(initializer)
1100            else:
1101                weights_to_keep.append(initializer.name)
1102        for initializer in weights_to_remove:
1103            graph.initializer.remove(initializer)
1104
1105        names_to_remove = [initializer.name for initializer in weights_to_remove]
1106        logger.debug(f"remove {len(weights_to_remove)} unused initializers: {names_to_remove}")
1107        if verbose:
1108            logger.debug(f"remaining initializers:{weights_to_keep}")
1109
1110        self.remove_unused_constant()
1111
1112    def is_safe_to_fuse_nodes(self, nodes_to_remove, keep_outputs, input_name_to_nodes, output_name_to_node):
1113        for node_to_remove in nodes_to_remove:
1114            for output_to_remove in node_to_remove.output:
1115                if output_to_remove in keep_outputs:
1116                    continue
1117
1118                if output_to_remove in input_name_to_nodes:
1119                    for impacted_node in input_name_to_nodes[output_to_remove]:
1120                        if impacted_node not in nodes_to_remove:
1121                            logger.debug(
1122                                "it is not safe to remove nodes since output %s is used by %s",
1123                                output_to_remove,
1124                                impacted_node,
1125                            )
1126                            return False
1127        return True
1128
1129    @staticmethod
1130    def graph_topological_sort(graph, is_deterministic=False):
1131        deps_set = set()  # dependency set of all node
1132        sorted_node_set = set()  # sorted node set
1133        sorted_nodes = []  # initialize sorted_nodes
1134
1135        initializer_names = [init.name for init in graph.initializer]
1136        graph_input_names = [input.name for input in graph.input]
1137        input_names = initializer_names + graph_input_names
1138
1139        if is_deterministic:
1140            input_names.sort()
1141
1142        for input_name in input_names:
1143            deps_set.add(input_name)
1144
1145        sorted_node_set_len = -1
1146        graph_nodes = graph.node if not is_deterministic else sorted(graph.node, key=lambda x: x.name)
1147
1148        last_node_name = None
1149        while len(sorted_node_set) != len(graph_nodes):
1150            if len(sorted_node_set) == sorted_node_set_len:
1151                break
1152            sorted_node_set_len = len(sorted_node_set)
1153            for node_idx, node in enumerate(graph_nodes):
1154                if node_idx in sorted_node_set:
1155                    continue
1156                input_count = sum(1 for _ in node.input if _)
1157                if input_count == 0:
1158                    sorted_nodes.append(node)
1159                    sorted_node_set.add(node_idx)
1160                    for output in node.output:
1161                        if output:
1162                            deps_set.add(output)
1163                    continue
1164                failed = False
1165                for input_name in node.input:
1166                    if input_name and input_name not in deps_set:
1167                        failed = True
1168                        last_node_name = node.name
1169                if not failed:
1170                    sorted_nodes.append(node)
1171                    sorted_node_set.add(node_idx)
1172                    for output in node.output:
1173                        if output:
1174                            deps_set.add(output)
1175                else:
1176                    continue
1177
1178        if len(sorted_node_set) != len(graph.node):
1179            raise RuntimeError(
1180                f"Graph is not a DAG: len(sorted_node_set)={len(sorted_node_set)}, len(graph.node)={len(graph.node)}, failed at node {last_node_name}"
1181            )
1182
1183        graph.ClearField("node")
1184        graph.node.extend(sorted_nodes)
1185
1186    def topological_sort(self, is_deterministic=False, dump_model_on_failure=False):
1187        # TODO: support graph_topological_sort() in subgraphs
1188        # for graph in self.graphs():
1189        #    self.graph_topological_sort(graph)
1190        try:
1191            OnnxModel.graph_topological_sort(self.model.graph, is_deterministic)
1192        except RuntimeError as e:
1193            if dump_model_on_failure:
1194                logger.info(
1195                    "Failed to sort graph in topological order. Dumping model to _topo_sort_failed.onnx for debugging."
1196                )
1197                OnnxModel.save(
1198                    self.model, "_topo_sort_failed.onnx", save_as_external_data=True, all_tensors_to_one_file=True
1199                )
1200            raise e

Showing the first 1,200 of 1637 lines. Download the file for the rest.

codekingpro/portable-devtools · Team Ai