codekingpro/portable-devtools
114k
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
